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"],