File size: 5,281 Bytes
fdf01c1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
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])