[model] add GLM-5.3-Flash training support (#10840)

This commit is contained in:
xvxuopop
2026-09-28 16:30:40 +08:00
committed by GitHub
parent 4d6c7cf03b
commit c1dc0f51ab
4 changed files with 350 additions and 1 deletions

View File

@@ -19,7 +19,7 @@ import inspect
import math
import os
import re
from copy import deepcopy
from copy import copy, deepcopy
from dataclasses import dataclass
from io import BytesIO
from types import SimpleNamespace
@@ -56,6 +56,7 @@ if TYPE_CHECKING:
from transformers.feature_extraction_sequence_utils import SequenceFeatureExtractor
from transformers.image_processing_utils import BaseImageProcessor
from transformers.video_processing_utils import BaseVideoProcessor
from transformers.video_utils import VideoMetadata
class EncodedImage(TypedDict):
path: str | None
@@ -2776,6 +2777,213 @@ class Qwen3VLPlugin(Qwen2VLPlugin):
return messages
@dataclass
class Glm5NextPlugin(BasePlugin):
r"""GLM5-Next images and timestamped video frames using the native processors."""
@override
def _validate_input(
self,
processor: Optional["MMProcessor"],
images: list["ImageInput"],
videos: list["VideoInput"],
audios: list["AudioInput"],
) -> None:
if processor is None and not (images or videos or audios):
return # Preserve tokenizer-only text preprocessing.
super()._validate_input(processor, images, videos, audios)
def _decode_video(
self, video: "VideoInput", processor: "MMProcessor"
) -> tuple[list["ImageObject"], "VideoMetadata"]:
r"""Decode sampled frames and preserve their source timing for the native processor."""
from transformers.video_utils import VideoMetadata
fps = getattr(processor, "video_fps", 2.0)
maxlen = getattr(processor, "video_maxlen", 128)
temporal = processor.video_processor.temporal_patch_size
if not math.isfinite(fps) or fps <= 0 or maxlen < temporal:
raise ValueError("glm5_next requires video_fps > 0 and video_maxlen >= temporal_patch_size.")
# Native sampling can duplicate the last frame for temporal patching.
# Use a private sampler so frame limits never mutate the shared processor.
sampler = copy(processor.video_processor)
sampler.max_frames = min(sampler.max_frames, maxlen // temporal * temporal)
if _check_video_is_nested_images(video):
if not video:
raise ValueError("glm5_next received an empty video frame list.")
total = len(video)
indices = np.linspace(0, total - 1, min(total, sampler.max_frames), dtype=int).tolist()
frames = self._regularize_images(
[video[i] for i in indices], image_max_pixels=float("inf"), image_min_pixels=0
)["images"]
# Frame lists have no source timing; video_fps describes their spacing.
metadata = VideoMetadata(total_num_frames=total, fps=fps, frames_indices=indices)
else:
with av.open(video, "r") as container:
stream = next((s for s in container.streams if s.type == "video"), None)
if stream is None or not stream.average_rate or stream.average_rate <= 0:
raise ValueError("glm5_next requires a video stream with a valid source FPS.")
source_fps = float(stream.average_rate)
duration = float(stream.duration * stream.time_base) if stream.duration is not None else None
metadata = VideoMetadata(total_num_frames=stream.frames, fps=source_fps, duration=duration)
# Some containers omit frame counts. Count without retaining all
# decoded pixels, then seek back and retain only selected frames.
if not metadata.total_num_frames:
metadata.total_num_frames = sum(1 for _ in container.decode(stream))
container.seek(0)
if metadata.total_num_frames <= 0:
raise ValueError("glm5_next received a video with no decodable frames.")
indices = sampler.sample_frames(metadata, fps=fps).tolist()
if not indices: # Sub-second clips can round to zero in the native sampler.
indices = [0]
indices = indices[: sampler.max_frames]
selected = set(indices)
decoded = {}
for i, frame in enumerate(container.decode(stream)):
if i in selected:
decoded[i] = frame.to_image()
if i >= max(selected):
break
if selected - decoded.keys():
raise ValueError("glm5_next could not decode all sampled video frames.")
frames = [decoded[i] for i in indices]
metadata.frames_indices = indices
# Keep both pixels and timestamps aligned when repeating an odd last frame.
while len(frames) % temporal:
frames.append(frames[-1])
metadata.frames_indices.append(metadata.frames_indices[-1])
return frames, metadata
@override
def _get_mm_inputs(
self,
images: list["ImageInput"],
videos: list["VideoInput"],
audios: list["AudioInput"],
processor: "MMProcessor",
) -> dict[str, Any]:
self._validate_input(processor, images, videos, audios)
mm_inputs = {}
if images:
# Only decode/convert here: resizing twice changes the native pixels.
images = self._regularize_images(images, image_max_pixels=float("inf"), image_min_pixels=0)["images"]
image_processor = processor.image_processor
pixels_per_token = (image_processor.patch_size * image_processor.merge_size) ** 2
kwargs = {}
if hasattr(processor, "image_max_pixels"):
kwargs["max_image_tokens"] = max(1, processor.image_max_pixels // pixels_per_token)
if hasattr(processor, "image_min_pixels"):
kwargs["min_image_tokens"] = max(1, math.ceil(processor.image_min_pixels / pixels_per_token))
mm_inputs.update(image_processor(images=images, return_tensors="pt", **kwargs))
if videos:
video_processor = processor.video_processor
pixels_per_token = (video_processor.patch_size * video_processor.merge_size) ** 2
processed = []
for video in videos:
frames, metadata = self._decode_video(video, processor)
temporal_tokens = len(frames) // video_processor.temporal_patch_size
# LlamaFactory limits pixels per frame; native GLM5 budgets tokens
# across the whole video (after temporal merging).
kwargs = {
"min_image_tokens": max(
1, math.ceil(getattr(processor, "video_min_pixels", 256) / pixels_per_token)
)
* temporal_tokens,
"max_image_tokens": max(1, getattr(processor, "video_max_pixels", 65536) // pixels_per_token)
* temporal_tokens,
}
processed.append(
video_processor(
videos=[np.stack([np.asarray(frame) for frame in frames])],
video_metadata=[metadata],
do_sample_frames=False,
return_metadata=True,
return_tensors="pt",
**kwargs,
)
)
for key in ("pixel_values_videos", "video_grid_thw"):
mm_inputs[key] = torch.cat([item[key] for item in processed])
mm_inputs["video_metadata"] = [item["video_metadata"][0] for item in processed]
return mm_inputs
@override
def process_messages(
self,
messages: list[dict[str, str]],
images: list["ImageInput"],
videos: list["VideoInput"],
audios: list["AudioInput"],
processor: Optional["MMProcessor"],
) -> list[dict[str, str]]:
self._validate_input(processor, images, videos, audios)
self._validate_messages(messages, images, videos, audios)
messages = deepcopy(messages)
mm_inputs = self._get_mm_inputs(images, videos, audios, processor) if self.expand_mm_tokens else {}
image_idx, video_idx = 0, 0
for message in messages:
if message["role"] != "user" and any(
p in message["content"] for p in (IMAGE_PLACEHOLDER, VIDEO_PLACEHOLDER)
):
raise ValueError("glm5_next supports images and videos in user messages only.")
while IMAGE_PLACEHOLDER in message["content"]:
tokens = (
processor.replace_image_token(mm_inputs, image_idx) if self.expand_mm_tokens else self.image_token
)
message["content"] = message["content"].replace(
IMAGE_PLACEHOLDER, f"<|begin_of_image|>{tokens}<|end_of_image|>", 1
)
image_idx += 1
while VIDEO_PLACEHOLDER in message["content"]:
tokens = (
processor.replace_video_token(mm_inputs, video_idx) if self.expand_mm_tokens else self.video_token
)
message["content"] = message["content"].replace(
VIDEO_PLACEHOLDER, f"<|begin_of_video|>{tokens}<|end_of_video|>", 1
)
video_idx += 1
return messages
@override
def get_mm_inputs(
self,
images: list["ImageInput"],
videos: list["VideoInput"],
audios: list["AudioInput"],
imglens: list[int],
vidlens: list[int],
audlens: list[int],
batch_ids: list[list[int]],
processor: Optional["MMProcessor"],
) -> dict[str, Union[list[list[int]], "torch.Tensor"]]:
mm_inputs = self._get_mm_inputs(images, videos, audios, processor)
token_types = processor.create_mm_token_type_ids(batch_ids)
for count, ids in zip(vidlens, batch_ids):
boundaries = [token for token in ids if token in (processor.video_start_id, processor.video_end_id)]
if boundaries != [processor.video_start_id, processor.video_end_id] * count:
raise ValueError("glm5_next video boundaries do not match inputs; check cutoff_len.")
# Images and video frames share image_token_id: validate by modality,
# otherwise a truncated image could be hidden by extra video tokens.
for modality, lens, key, subprocessor in (
(1, imglens, "image_grid_thw", processor.image_processor),
(2, vidlens, "video_grid_thw", processor.video_processor),
):
grids = mm_inputs.get(key, [])
offset = 0
for count, types in zip(lens, token_types):
expected = sum(
int(grid.prod()) // subprocessor.merge_size**2 for grid in grids[offset : offset + count]
)
if types.count(modality) != expected:
raise ValueError(
"glm5_next image/video tokens do not match features; check cutoff_len and media limits."
)
offset += count
mm_inputs.pop("video_metadata", None) # Used for prompt timestamps, never forwarded to model.forward.
mm_inputs["mm_token_type_ids"] = token_types
return mm_inputs
@dataclass
class GLM4VPlugin(Qwen2VLPlugin):
@override
@@ -3249,6 +3457,7 @@ PLUGINS = {
"gemma3n": Gemma3nPlugin,
"gemma4": Gemma4Plugin,
"glm4v": GLM4VPlugin,
"glm5_next": Glm5NextPlugin,
"intern_vl": InternVLPlugin,
"kimi_vl": KimiVLPlugin,
"llama4": Llama4Plugin,

View File

@@ -339,6 +339,73 @@ class Template:
return modelfile
@dataclass
class Glm5NextTemplate(Template):
r"""GLM-5.3-Flash template with preserved thinking and EOS only on the final assistant response."""
@override
def _encode(
self,
tokenizer: "PreTrainedTokenizer",
messages: list[dict[str, str]],
system: Optional[str],
tools: Optional[str],
) -> list[list[int]]:
if not messages or len(messages) % 2 or messages[0]["role"] != Role.USER:
raise ValueError(
"glm5_next expects alternating user/observation and assistant/function pairs, starting with user."
)
system = system or self.default_system
encoded = []
for i, message in enumerate(messages):
role, content = message["role"], message["content"]
expected = (Role.USER, Role.OBSERVATION) if i % 2 == 0 else (Role.ASSISTANT, Role.FUNCTION)
if role not in expected or not isinstance(content, str):
raise ValueError(
"glm5_next expects alternating user/observation and assistant/function text messages."
)
if message.get("tool_calls"):
raise ValueError("glm5_next expects tool calls as JSON in role=function content (v0 ShareGPT format).")
if i % 2 == 0 and i > 0:
after_function = messages[i - 1]["role"] == Role.FUNCTION
if (role == Role.OBSERVATION) != after_function:
raise ValueError("glm5_next requires an observation after each function turn.")
elements = []
if i == 0:
elements += self.format_prefix.apply()
if tools:
elements += self.format_tools.apply(content=tools)
if system:
elements += self.format_system.apply(content=system)
if role == Role.USER:
elements += self.format_user.apply(content=content)
elif role == Role.OBSERVATION:
elements += self.format_observation.apply(content=content)
else:
if role == Role.FUNCTION:
content = self.format_function.apply(
content=content,
thought_words=self.thought_words,
tool_call_words=self.tool_call_words,
)[0]
if "</think>" in content:
reasoning = content.split("</think>")[0].split("<think>")[-1]
content = content.split("</think>")[-1]
else:
reasoning = ""
elements += self.format_assistant.apply(content=reasoning + "</think>" + content.strip())
if role == Role.FUNCTION:
# Supervise the observation marker so the model learns to hand control to tools.
elements += ["<|observation|>"]
elif i == len(messages) - 1:
elements += [{"eos_token"}]
encoded.append(self._convert_elements_to_ids(tokenizer, elements))
return encoded
@dataclass
class MossVLTemplate(Template):
@override
@@ -1211,6 +1278,23 @@ register_template(
)
register_template(
name="glm5_next",
format_user=StringFormatter(slots=["<|user|>{{content}}<|assistant|><think>"]),
format_assistant=StringFormatter(slots=["{{content}}"]),
format_system=StringFormatter(slots=["<|system|>{{content}}"]),
format_function=FunctionFormatter(slots=["{{content}}"], tool_format="glm5_next"),
format_observation=StringFormatter(slots=["<tool_response>{{content}}</tool_response><|assistant|><think>"]),
format_tools=ToolFormatter(tool_format="glm5_next"),
format_prefix=EmptyFormatter(slots=["[gMASK]<sop><|system|>Reasoning Effort: Max"]),
stop_words=["<|user|>", "<|observation|>"],
thought_words=("<think>", "</think>"),
preserve_thinking=True,
mm_plugin=get_mm_plugin(name="glm5_next", image_token="<|image|>", video_token="<|video|>"),
template_class=Glm5NextTemplate,
)
register_template(
name="glm4",
format_user=StringFormatter(slots=["<|user|>\n{{content}}<|assistant|>"]),

View File

@@ -55,6 +55,14 @@ GLM4_MOE_TOOL_PROMPT = (
"\n...\n</tool_call>\n"
)
GLM5_NEXT_TOOL_PROMPT = (
"<|system|>\n# Tools\n\nYou may call one or more functions to assist with the user query.\n\n"
"You are provided with function signatures within <tools></tools> XML tags:\n<tools>\n{tool_text}"
"</tools>\n\nFor each function call, output the function name and arguments within the following XML format:\n"
"<tool_call>{{function-name}}<arg_key>{{arg-key-1}}</arg_key><arg_value>{{arg-value-1}}</arg_value>"
"<arg_key>{{arg-key-2}}</arg_key><arg_value>{{arg-value-2}}</arg_value>...</tool_call>"
)
LLAMA3_TOOL_PROMPT = (
"Cutting Knowledge Date: December 2023\nToday Date: {date}\n\n"
"You have access to the following functions. To call a function, please respond with JSON for a function call. "
@@ -804,6 +812,44 @@ class GLM4MOEToolUtils(QwenToolUtils):
return "\n".join(function_texts)
class GLM5NextToolUtils(GLM4MOEToolUtils):
r"""GLM5-Next tool using template."""
@override
@staticmethod
def tool_formatter(tools: list[dict[str, Any]]) -> str:
if not isinstance(tools, list):
raise ValueError("glm5_next tools must be a JSON list.")
tool_text = ""
for tool in tools:
tool = tool.get("function", tool)
if tool.get("defer_loading", False):
continue
tool = {key: value for key, value in tool.items() if key not in {"strict", "defer_loading"}}
tool_text += json.dumps(tool, ensure_ascii=False) + "\n"
return GLM5_NEXT_TOOL_PROMPT.format(tool_text=tool_text)
@override
@staticmethod
def function_formatter(functions: list["FunctionCall"]) -> str:
if not functions:
raise ValueError("glm5_next function messages must contain at least one call.")
calls = []
for name, arguments in functions:
if not isinstance(name, str) or not re.fullmatch(r"[^\s<>]+", name):
raise ValueError("Invalid glm5_next function name.")
arguments = json.loads(arguments)
if not isinstance(arguments, dict):
raise ValueError("glm5_next function arguments must be a JSON object.")
text = "<tool_call>" + name
for key, value in arguments.items():
value = value if isinstance(value, str) else json.dumps(value, ensure_ascii=False)
text += f"<arg_key>{key}</arg_key><arg_value>{value}</arg_value>"
calls.append(text + "</tool_call>")
return "".join(calls)
class SeedToolUtils(ToolUtils):
r"""Seed tool using template."""
@@ -986,6 +1032,7 @@ TOOLS = {
"qwen3_5": Qwen35ToolUtils(),
"qwen3_8": Qwen38ToolUtils(),
"glm4_moe": GLM4MOEToolUtils(),
"glm5_next": GLM5NextToolUtils(),
"seed_oss": SeedToolUtils(),
"ling": LingToolUtils(),
}

View File

@@ -261,6 +261,15 @@ _register_composite_model(
)
_register_composite_model(
model_type="glm5_next",
projector_keys=["model.visual.merger", "model.visual.downsample"],
vision_model_keys=["model.visual.patch_embed", "model.visual.blocks", "model.visual.post_layernorm"],
language_model_keys=["model.language_model", "lm_head"],
lora_conflict_keys=["patch_embed"],
)
_register_composite_model(
model_type="glm_ocr",
projector_keys=["visual.merger"],