File size: 4,081 Bytes
6a807da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ea5c483
 
6a807da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
103
104
105
106
107
108
109
110
111
112
# ensemble_qwen.py
import torch, math
from torch import nn
from transformers import Qwen3ForCausalLM, Qwen3Config, PreTrainedModel, PretrainedConfig


from transformers.modeling_outputs import CausalLMOutputWithPast
from transformers.generation import GenerationMixin

class EnsembleConfig(Qwen3Config):
    model_type = "ensemble_qwen"
    def __init__(
        self,
        model_a_path="Qwen/Qwen3-4B",
        model_b_path="Qwen/Qwen3-4B-Base",
        weight_a=0.9,
        weight_b=0.1,
        **kwargs,
    ):
        super().__init__(**kwargs)
        self.model_a_path, self.model_b_path = model_a_path, model_b_path
        self.weight_a, self.weight_b = weight_a, weight_b


class EnsembleForCausalLM(PreTrainedModel, GenerationMixin):
    config_class = EnsembleConfig
    _supports_flash_attn_2 = True
    _supports_sdpa        = True
    main_input_name = "input_ids"
    _tp_plan = {"model_a.lm_head": "colwise_rep", "model_b.lm_head": "colwise_rep"}
    
    def __init__(self, config: EnsembleConfig):
        super().__init__(config)
        
        self.weight_a = config.weight_a
        self.weight_b = config.weight_b
        
        # remove extra keys from config and initialzie Qwen3 model
        qwen3_config_dict = config.to_dict()
        extra_keys = ["model_a_path", "model_b_path", "weight_a", "weight_b", "auto_map"]
        for key in extra_keys:
            del qwen3_config_dict[key]
        
        qwen3_config_dict["model_type"] = "qwen3"
        qwen3_config_dict["architectures"] = ["Qwen3ForCausalLM"]
        
        if hasattr(config, "attn_implementation"):
            qwen3_config_dict["attn_implementation"] = config.attn_implementation
            
        qwen3_config = Qwen3Config(**qwen3_config_dict)
        
        self.model_a = Qwen3ForCausalLM(qwen3_config)
        self.model_b = Qwen3ForCausalLM(qwen3_config)

    # ---- core magic ----------------------------------------------------
    def forward(
        self,
        input_ids: torch.LongTensor,
        attention_mask: torch.LongTensor | None = None,
        past_key_values: tuple | None = None,
        use_cache: bool = True,
        **kwargs,
    ) -> CausalLMOutputWithPast:
        """Run both models, sum logits, carry both kv-caches forward."""
        past_a, past_b = (None, None) if past_key_values is None else past_key_values
        
        out_a = self.model_a(
            input_ids,
            attention_mask=attention_mask,
            past_key_values=past_a,
            use_cache=use_cache,
        )
        out_b = self.model_b(
            input_ids,
            attention_mask=attention_mask,
            past_key_values=past_b,
            use_cache=use_cache,
        )

        # blend logits
        logits = self.weight_a * out_a.logits + self.weight_b * out_b.logits

        # pack both pasts together so generate() can recycle them next step
        past = (out_a.past_key_values, out_b.past_key_values) if use_cache else None

        return CausalLMOutputWithPast(
            logits=logits,
            past_key_values=past,
            hidden_states=None,
            attentions=None,
        )

    # ---- wiring for HF generation loop ---------------------------------
    def prepare_inputs_for_generation(self, input_ids, past_key_values=None, **kwargs):
        # HF expects past_key_values shape to match forward() signature
        if past_key_values is not None:
            past_a, past_b = past_key_values
            input_ids = input_ids[:, -1:]     # only last token in incremental mode
        else:
            past_a = past_b = None
        return {
            "input_ids": input_ids,
            "past_key_values": (past_a, past_b),
            "use_cache": kwargs.get("use_cache", True),
        }

    def _reorder_cache(self, past, beam_idx):
        # support beam search ↔ reorder both caches identically
        past_a, past_b = past
        past_a = self.model_a._reorder_cache(past_a, beam_idx)
        past_b = self.model_b._reorder_cache(past_b, beam_idx)
        return (past_a, past_b)