mirror of
https://github.com/hiyouga/LLaMA-Factory.git
synced 2026-10-05 06:15:44 +08:00
[model] add GLM-5.3-Flash training support (#10840)
This commit is contained in:
@@ -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,
|
||||||
|
|||||||
@@ -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|>"]),
|
||||||
|
|||||||
@@ -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(),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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"],
|
||||||
|
|||||||
Reference in New Issue
Block a user