diff --git a/src/llamafactory/data/mm_plugin.py b/src/llamafactory/data/mm_plugin.py
index 10eefd3db..a299c93f0 100644
--- a/src/llamafactory/data/mm_plugin.py
+++ b/src/llamafactory/data/mm_plugin.py
@@ -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,
diff --git a/src/llamafactory/data/template.py b/src/llamafactory/data/template.py
index 08af15464..d499b4e00 100644
--- a/src/llamafactory/data/template.py
+++ b/src/llamafactory/data/template.py
@@ -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 "" in content:
+ reasoning = content.split("")[0].split("")[-1]
+ content = content.split("")[-1]
+ else:
+ reasoning = ""
+ elements += self.format_assistant.apply(content=reasoning + "" + 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|>"]),
+ format_assistant=StringFormatter(slots=["{{content}}"]),
+ format_system=StringFormatter(slots=["<|system|>{{content}}"]),
+ format_function=FunctionFormatter(slots=["{{content}}"], tool_format="glm5_next"),
+ format_observation=StringFormatter(slots=["{{content}}<|assistant|>"]),
+ format_tools=ToolFormatter(tool_format="glm5_next"),
+ format_prefix=EmptyFormatter(slots=["[gMASK]<|system|>Reasoning Effort: Max"]),
+ stop_words=["<|user|>", "<|observation|>"],
+ thought_words=("", ""),
+ 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|>"]),
diff --git a/src/llamafactory/data/tool_utils.py b/src/llamafactory/data/tool_utils.py
index d2f322d32..b74fdc65c 100644
--- a/src/llamafactory/data/tool_utils.py
+++ b/src/llamafactory/data/tool_utils.py
@@ -55,6 +55,14 @@ GLM4_MOE_TOOL_PROMPT = (
"\n...\n\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 XML tags:\n\n{tool_text}"
+ "\n\nFor each function call, output the function name and arguments within the following XML format:\n"
+ "{{function-name}}{{arg-key-1}}{{arg-value-1}}"
+ "{{arg-key-2}}{{arg-value-2}}..."
+)
+
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 = "" + name
+ for key, value in arguments.items():
+ value = value if isinstance(value, str) else json.dumps(value, ensure_ascii=False)
+ text += f"{key}{value}"
+ calls.append(text + "")
+ 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(),
}
diff --git a/src/llamafactory/model/model_utils/visual.py b/src/llamafactory/model/model_utils/visual.py
index 9a5e80a98..98b712db1 100644
--- a/src/llamafactory/model/model_utils/visual.py
+++ b/src/llamafactory/model/model_utils/visual.py
@@ -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"],