Valen-4B / batching.py
laolao77's picture
Upload Valen-4B merged model and unmerged LoRA bundle
fdf01c1 verified
Raw History Blame Contribute Delete
2.95 kB
"""Pad language rows and concatenate visual patches. / 文本补齐,视觉 patch 拼接。"""
import torch
from torch.nn.utils.rnn import pad_sequence
SEQUENCE_KEYS = {"input_ids", "attention_mask", "mm_token_type_ids"}
VISUAL_KEYS = {"pixel_values", "pixel_values_videos", "image_grid_thw", "video_grid_thw"}
def collate_branches(branches, pad_token_id=0):
if not branches:
raise ValueError("Cannot collate an empty branch batch")
inputs = [branch.inputs for branch in branches]
unknown = set().union(*(row.keys() for row in inputs)) - SEQUENCE_KEYS - VISUAL_KEYS
if unknown:
raise ValueError(f"Unsupported batched Qwen inputs: {sorted(unknown)}")
if any(row["input_ids"].ndim != 2 or row["input_ids"].shape[0] != 1 for row in inputs):
raise ValueError("Each compiled branch must contain one language row")
result = {}
for key in SEQUENCE_KEYS:
if key != "input_ids" and not any(key in row for row in inputs):
continue
rows = []
for row in inputs:
value = row.get(key)
if value is None:
value = (torch.ones_like(row["input_ids"]) if key == "attention_mask"
else torch.zeros_like(row["input_ids"]))
if value.shape != row["input_ids"].shape:
raise ValueError(f"{key} must match the branch input_ids shape")
rows.append(value[0])
# Right padding keeps all compiler readout positions and spans unchanged.
# 右侧补齐,候选位置和角色区间无需偏移。
result[key] = pad_sequence(rows, batch_first=True,
padding_value=pad_token_id if key == "input_ids" else 0)
if "attention_mask" not in result:
result["attention_mask"] = pad_sequence([torch.ones_like(r["input_ids"][0]) for r in inputs],
batch_first=True, padding_value=0)
for key in VISUAL_KEYS:
values = [row[key] for row in inputs if key in row]
if values:
result[key] = torch.cat(values, dim=0)
return result
def branch_batches(questions, max_tokens):
"""Limit padded tokens per backbone forward, keeping Score branch order."""
if type(max_tokens) is not int or max_tokens <= 0:
raise ValueError("microbatch_max_tokens must be a positive integer")
batch, longest = [], 0
for question_index, question in enumerate(questions):
if not question.branches:
raise ValueError("A Qwen question requires at least one branch")
for branch in question.branches:
length = branch.inputs["input_ids"].shape[1]
if batch and max(longest, length) * (len(batch) + 1) > max_tokens:
yield batch
batch, longest = [], 0
batch.append((question_index, branch))
longest = max(longest, length)
if batch:
yield batch