[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 math
import os import os
import re import re
from copy import deepcopy from copy import copy, deepcopy
from dataclasses import dataclass from dataclasses import dataclass
from io import BytesIO from io import BytesIO
from types import SimpleNamespace from types import SimpleNamespace
@@ -56,6 +56,7 @@ if TYPE_CHECKING:
from transformers.feature_extraction_sequence_utils import SequenceFeatureExtractor from transformers.feature_extraction_sequence_utils import SequenceFeatureExtractor
from transformers.image_processing_utils import BaseImageProcessor from transformers.image_processing_utils import BaseImageProcessor
from transformers.video_processing_utils import BaseVideoProcessor from transformers.video_processing_utils import BaseVideoProcessor
from transformers.video_utils import VideoMetadata
class EncodedImage(TypedDict): class EncodedImage(TypedDict):
path: str | None path: str | None
@@ -2776,6 +2777,213 @@ class Qwen3VLPlugin(Qwen2VLPlugin):
return messages 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 @dataclass
class GLM4VPlugin(Qwen2VLPlugin): class GLM4VPlugin(Qwen2VLPlugin):
@override @override
@@ -3249,6 +3457,7 @@ PLUGINS = {
"gemma3n": Gemma3nPlugin, "gemma3n": Gemma3nPlugin,
"gemma4": Gemma4Plugin, "gemma4": Gemma4Plugin,
"glm4v": GLM4VPlugin, "glm4v": GLM4VPlugin,
"glm5_next": Glm5NextPlugin,
"intern_vl": InternVLPlugin, "intern_vl": InternVLPlugin,
"kimi_vl": KimiVLPlugin, "kimi_vl": KimiVLPlugin,
"llama4": Llama4Plugin, "llama4": Llama4Plugin,

View File

@@ -339,6 +339,73 @@ class Template:
return modelfile 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 @dataclass
class MossVLTemplate(Template): class MossVLTemplate(Template):
@override @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( register_template(
name="glm4", name="glm4",
format_user=StringFormatter(slots=["<|user|>\n{{content}}<|assistant|>"]), format_user=StringFormatter(slots=["<|user|>\n{{content}}<|assistant|>"]),

View File

@@ -55,6 +55,14 @@ GLM4_MOE_TOOL_PROMPT = (
"\n...\n</tool_call>\n" "\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 = ( LLAMA3_TOOL_PROMPT = (
"Cutting Knowledge Date: December 2023\nToday Date: {date}\n\n" "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. " "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) 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): class SeedToolUtils(ToolUtils):
r"""Seed tool using template.""" r"""Seed tool using template."""
@@ -986,6 +1032,7 @@ TOOLS = {
"qwen3_5": Qwen35ToolUtils(), "qwen3_5": Qwen35ToolUtils(),
"qwen3_8": Qwen38ToolUtils(), "qwen3_8": Qwen38ToolUtils(),
"glm4_moe": GLM4MOEToolUtils(), "glm4_moe": GLM4MOEToolUtils(),
"glm5_next": GLM5NextToolUtils(),
"seed_oss": SeedToolUtils(), "seed_oss": SeedToolUtils(),
"ling": LingToolUtils(), "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( _register_composite_model(
model_type="glm_ocr", model_type="glm_ocr",
projector_keys=["visual.merger"], projector_keys=["visual.merger"],