import torch from torch import nn MODEL_NAME = "Valen" from .heads import DecisionHead, build_head from .batching import branch_batches, collate_branches class ValenQwen(nn.Module): """Qwen3.5 backbone with a shared candidate decision head.""" model_name = MODEL_NAME architecture = "qwen" def __init__(self, backbone, projection_dim=256, head_config=None): super().__init__() self.backbone = backbone self.head = build_head(backbone.config.text_config.hidden_size, dict({"projection_dim": projection_dim}, **(head_config or {}))) def forward(self, question): device = next(self.backbone.parameters()).device outputs = [] for branch in question.branches: inputs = {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in branch.inputs.items()} hidden = self.backbone(**inputs, use_cache=False, return_dict=True).last_hidden_state outputs.append(self.head.score_features(self.head.extract_features(hidden, branch))) return torch.cat(outputs) def extract_features(self, question): """缓存投影前的特征。 / Cache raw readouts before trainable head layers.""" device = next(self.backbone.parameters()).device features = [] for branch in question.branches: inputs = {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in branch.inputs.items()} hidden = self.backbone(**inputs, use_cache=False, return_dict=True).last_hidden_state features.append(self.head.extract_features(hidden, branch)) return features def extract_state_features(self, state): """一次编码 state,读取各题特征。 / One backbone graph for every question.""" if not state.questions: return [] if state.inputs is None: raise ValueError("Shared-state features require compiled shared inputs") device = next(self.backbone.parameters()).device inputs = {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in state.inputs.items()} hidden = self.backbone(**inputs, use_cache=False, return_dict=True).last_hidden_state return self._state_features(hidden, state) def _state_features(self, hidden, state): return [[self.head.extract_features(hidden, branch) for branch in q.branches] for q in state.questions] def forward_state(self, state): return [self.score_features(features) for features in self.extract_state_features(state)] def forward_state_batch(self, states, max_tokens=32768): """不同 state 组成 batch,同一 state 的题目共用一行。 / One row per shared state.""" device = next(self.backbone.parameters()).device pad_token_id = getattr(self.backbone.config.text_config, "pad_token_id", None) or 0 outputs = [[] for _ in states] representatives = [] indices = [] for index, state in enumerate(states): if not state.questions: continue if state.inputs is None or any(branch.inputs is not state.inputs for q in state.questions for branch in q.branches): raise ValueError("Shared-state batch requires one shared input per state") representatives.append(state.questions[0]) indices.append(index) for batch in branch_batches(representatives, max_tokens): inputs = collate_branches([branch for _, branch in batch], pad_token_id) inputs = {key: value.to(device) for key, value in inputs.items()} hidden = self.backbone(**inputs, use_cache=False, return_dict=True).last_hidden_state for row, (index, _) in enumerate(batch): original = indices[index] features = self._state_features(hidden[row:row + 1], states[original]) outputs[original] = [self.score_features(f) for f in features] return outputs def forward_batch(self, questions, max_tokens=32768): """Forward several QAs together; each keeps its own decision head inputs. 多条 QA 并行;每条保留独立的候选、角色区间和 Score 分支。 """ device = next(self.backbone.parameters()).device pad_token_id = getattr(self.backbone.config.text_config, "pad_token_id", None) or 0 outputs = [[] for _ in questions] for batch in branch_batches(questions, max_tokens): inputs = collate_branches([branch for _, branch in batch], pad_token_id) inputs = {key: value.to(device) for key, value in inputs.items()} hidden = self.backbone(**inputs, use_cache=False, return_dict=True).last_hidden_state for row, (question_index, branch) in enumerate(batch): features = self.head.extract_features(hidden[row:row + 1], branch) outputs[question_index].append(self.head.score_features(features)) return [torch.cat(branches) for branches in outputs] def score_features(self, features, head=None): head = self.head if head is None else head return torch.cat([head.score_features(hidden) for hidden in features])