File size: 9,727 Bytes
00dd625
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
"""AlignX Model — full two-stage alignment framework.

Wraps a HuggingFace causal LM and surgically replaces the last transformer
FFN layer with the MoCaE AlignX layer (Stage 2 injection).
"""

import os
import torch
import torch.nn as nn
from typing import Optional, Dict, List, Union
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig

from .mocae import MoCaE, AlignXLayer, ExpertFFN
from .task_feature_matrix import TaskFeatureMatrix


# ---------------------------------------------------------------------------
# Model-family helpers
# ---------------------------------------------------------------------------

def _get_layers(model) -> nn.ModuleList:
    """Return the transformer layer list for supported model families."""
    if hasattr(model, "model") and hasattr(model.model, "layers"):
        return model.model.layers
    if hasattr(model, "transformer") and hasattr(model.transformer, "h"):
        return model.transformer.h
    raise ValueError(f"Unsupported model architecture: {type(model)}")


def _get_mlp(layer) -> nn.Module:
    """Return the FFN/MLP sub-module from a transformer layer."""
    for attr in ("mlp", "feed_forward", "ffn", "ff"):
        if hasattr(layer, attr):
            return getattr(layer, attr)
    raise ValueError(f"Cannot locate MLP in layer {type(layer)}")


def _set_mlp(layer, new_module: nn.Module):
    """Replace the FFN/MLP sub-module in a transformer layer."""
    for attr in ("mlp", "feed_forward", "ffn", "ff"):
        if hasattr(layer, attr):
            setattr(layer, attr, new_module)
            return
    raise ValueError(f"Cannot locate MLP in layer {type(layer)}")


def _get_intermediate_dim(mlp_module) -> int:
    """Infer the FFN intermediate dimension from the gate/up projection."""
    for attr in ("gate_proj", "w1", "fc1", "dense_h_to_4h"):
        if hasattr(mlp_module, attr):
            proj = getattr(mlp_module, attr)
            return proj.out_features
    for attr in ("up_proj", "w2", "fc2"):
        if hasattr(mlp_module, attr):
            return getattr(mlp_module, attr).out_features
    raise ValueError(f"Cannot infer intermediate_dim from {type(mlp_module)}")


# ---------------------------------------------------------------------------
# AlignX model class
# ---------------------------------------------------------------------------

class AlignXModel(nn.Module):
    """AlignX-wrapped causal LM with MoCaE layer injected at the last FFN.

    Usage:
        model = AlignXModel(base_lm, hidden_dim=4096, intermediate_dim=11008)
        model.register_task_matrices(T_h, T_ha, T_ho)
        # Train only MoCaE params (base_lm is frozen)
        logits = model(input_ids, attention_mask)
    """

    def __init__(
        self,
        base_lm: nn.Module,
        hidden_dim: int,
        intermediate_dim: int,
        k: int = 256,
        lambda1: float = 0.6,
        lambda2: float = 0.4,
        epsilon: float = 0.05,
        n_clusters: int = 8,
        layer_idx: int = -1,
        freeze_base: bool = True,
    ):
        super().__init__()
        self.base_lm = base_lm
        self.hidden_dim = hidden_dim
        self.layer_idx = layer_idx
        self.config = base_lm.config

        # Build MoCaE
        self.mocae = MoCaE(
            hidden_dim=hidden_dim,
            intermediate_dim=intermediate_dim,
            k=k,
            lambda1=lambda1,
            lambda2=lambda2,
            epsilon=epsilon,
            n_clusters=n_clusters,
        )
        self.alignx_layer = AlignXLayer(self.mocae, hidden_dim)

        # Inject AlignX layer into the last transformer FFN
        layers = _get_layers(self.base_lm)
        target_layer = layers[layer_idx]
        original_mlp = _get_mlp(target_layer)

        # Initialise all experts from the original FFN weights (helpful default)
        self.mocae.init_experts_from_ffn(original_mlp)

        # Replace the FFN
        _set_mlp(target_layer, self.alignx_layer)

        # Move AlignXLayer (MoCaE) to the same device as the rest of the target layer.
        # input_layernorm is a plain fp32/bf16 LayerNorm — its device is reliable even
        # when the base model is 4-bit quantized.
        if hasattr(target_layer, "input_layernorm"):
            _target_device = target_layer.input_layernorm.weight.device
        elif hasattr(target_layer, "ln_1"):
            _target_device = target_layer.ln_1.weight.device
        else:
            _target_device = torch.device(
                f"cuda:{torch.cuda.device_count() - 1}" if torch.cuda.is_available() else "cpu"
            )
        _target_dtype = target_layer.input_layernorm.weight.dtype if hasattr(target_layer, "input_layernorm") else torch.float16
        self.alignx_layer.to(_target_device).to(_target_dtype)
        print(f"[AlignX] AlignXLayer (MoCaE) placed on {_target_device} dtype={_target_dtype}")

        # Freeze base model parameters (only MoCaE is trainable)
        if freeze_base:
            for name, param in self.base_lm.named_parameters():
                param.requires_grad_(False)
        for param in self.mocae.parameters():
            param.requires_grad_(True)

        total_base = sum(p.numel() for p in self.base_lm.parameters())
        total_mocae = sum(p.numel() for p in self.mocae.parameters())
        print(f"[AlignX] Base LM params: {total_base:,}  |  MoCaE params: {total_mocae:,}")

    # ------------------------------------------------------------------

    def register_task_matrices(self, T_helpful, T_harmless, T_honest):
        self.mocae.register_task_matrices(T_helpful, T_harmless, T_honest)

    def init_expert_from_finetuned(self, expert_idx: int, finetuned_lm: nn.Module):
        """Initialise expert `expert_idx` from the last FFN of a fine-tuned model."""
        layers = _get_layers(finetuned_lm)
        ffn = _get_mlp(layers[self.layer_idx])
        self.mocae.init_expert_from_ffn(expert_idx, ffn)

    # ------------------------------------------------------------------

    def forward(self, input_ids, attention_mask=None, labels=None, **kwargs):
        return self.base_lm(
            input_ids=input_ids,
            attention_mask=attention_mask,
            labels=labels,
            **kwargs,
        )

    @torch.no_grad()
    def generate(self, input_ids, attention_mask=None, **kwargs):
        return self.base_lm.generate(
            input_ids=input_ids,
            attention_mask=attention_mask,
            **kwargs,
        )

    def save_mocae(self, path: str):
        os.makedirs(path, exist_ok=True)
        torch.save(self.mocae.state_dict(), os.path.join(path, "mocae.pt"))
        print(f"[AlignX] Saved MoCaE weights to {path}/mocae.pt")

    def load_mocae(self, path: str):
        state = torch.load(os.path.join(path, "mocae.pt"), map_location="cpu")
        self.mocae.load_state_dict(state, strict=False)
        print(f"[AlignX] Loaded MoCaE weights from {path}/mocae.pt")


# ---------------------------------------------------------------------------
# Builder helper
# ---------------------------------------------------------------------------

def build_alignx_model(
    base_model_name_or_path: str,
    finetuned_paths: Optional[Dict[str, str]] = None,
    task_matrix_paths: Optional[Dict[str, str]] = None,
    load_in_4bit: bool = True,
    device_map: str = "auto",
    k: int = 256,
    lambda1: float = 0.6,
    lambda2: float = 0.4,
    freeze_base: bool = True,
    layer_idx: int = -1,
) -> AlignXModel:
    """Build a full AlignX model from a base checkpoint.

    Optionally initialises each expert from a per-axis fine-tuned checkpoint
    and loads precomputed task-feature matrices.
    """
    bnb_config = None
    if load_in_4bit:
        bnb_config = BitsAndBytesConfig(
            load_in_4bit=True,
            bnb_4bit_compute_dtype=torch.bfloat16,
            bnb_4bit_use_double_quant=True,
            bnb_4bit_quant_type="nf4",
        )

    base_lm = AutoModelForCausalLM.from_pretrained(
        base_model_name_or_path,
        quantization_config=bnb_config,
        device_map=device_map,
        torch_dtype=torch.bfloat16,
        trust_remote_code=True,
    )

    hidden_dim = base_lm.config.hidden_size
    layers = _get_layers(base_lm)
    mlp = _get_mlp(layers[layer_idx])
    intermediate_dim = _get_intermediate_dim(mlp)

    alignx = AlignXModel(
        base_lm=base_lm,
        hidden_dim=hidden_dim,
        intermediate_dim=intermediate_dim,
        k=k,
        lambda1=lambda1,
        lambda2=lambda2,
        layer_idx=layer_idx,
        freeze_base=freeze_base,
    )

    # Initialise each expert from axis-specific fine-tuned model
    if finetuned_paths:
        axis_to_idx = {"helpful": 0, "harmless": 1, "honest": 2}
        for axis, ckpt_path in finetuned_paths.items():
            if axis in axis_to_idx and os.path.exists(ckpt_path):
                print(f"[AlignX] Loading expert {axis} from {ckpt_path}")
                ft_lm = AutoModelForCausalLM.from_pretrained(
                    ckpt_path,
                    torch_dtype=torch.float16,
                    device_map="cpu",
                    trust_remote_code=True,
                )
                alignx.init_expert_from_finetuned(axis_to_idx[axis], ft_lm)
                del ft_lm

    # Load task-feature matrices
    if task_matrix_paths:
        T = {}
        for axis in ("helpful", "harmless", "honest"):
            path = task_matrix_paths.get(axis, "")
            if path and os.path.exists(path):
                T[axis] = torch.load(path, map_location="cpu")
        if len(T) == 3:
            alignx.register_task_matrices(T["helpful"], T["harmless"], T["honest"])

    return alignx