File size: 12,528 Bytes
43e5946
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Backbone (Ministral-3-8B, Mistral3ForConditionalGeneration) + LoRA + joint schema head (DESIGN.md 2.1, 3.1).

- The inner model (`.model`, Mistral3Model) is called; its `last_hidden_state` is the post-norm final state.
  `lm_head` is never called and is deleted after load (131,072 x 4,096 bf16 = 1.07e9 B).
- Extra taps: outputs of decoder layers by index, captured with forward hooks.
- LoRA: language-model linear layers only (q/k/v/o/gate/up/down), vision tower + projector frozen, alpha = 2r.
  Same leaf discovery and vision-exclusion rule as the Standard One LoRA trainer.
"""
import json
import os
from pathlib import Path

import torch
import torch.nn as nn

from head import HeadConfig, JointSchemaHead
from render import token_roles

VISION_EXCLUDE_REGEX = r".*\.(vision_tower|multi_modal_projector)\..*"


def lora_leaves(model):
    """Leaf names of the language-model attention + MLP linears (as the trainer's `attention_mlp` targets)."""
    attention, mlp = set(), set()
    for name, module in model.named_modules():
        if not isinstance(module, nn.Linear) or ".language_model." not in name:
            continue
        leaf = name.rsplit(".", 1)[-1]
        if ".self_attn." in name:
            attention.add(leaf)
        elif ".mlp." in name and leaf in {"gate_proj", "up_proj", "down_proj"}:
            mlp.add(leaf)
    if not attention:
        raise ValueError("no attention projections found")
    return sorted(attention | mlp)


def decoder_layers(model):
    lm = model.model.language_model
    return lm.layers


def tiny_config(vocab_size=131072, layers=3, hidden=64, heads=4, kv_heads=2, intermediate=128):
    """A tiny, randomly initialised Mistral3 (Ministral3 text + Pixtral vision) config for CPU smoke tests."""
    from transformers import Mistral3Config
    text = {"model_type": "ministral3", "hidden_size": hidden, "intermediate_size": intermediate,
            "num_hidden_layers": layers, "num_attention_heads": heads, "num_key_value_heads": kv_heads,
            "head_dim": hidden // heads, "vocab_size": vocab_size, "max_position_embeddings": 262144,
            "rms_norm_eps": 1e-5, "hidden_act": "silu", "tie_word_embeddings": False, "sliding_window": None,
            "rope_parameters": {"beta_fast": 32.0, "beta_slow": 1.0, "factor": 16.0, "llama_4_scaling_beta": 0.1,
                                "mscale": 1.0, "mscale_all_dim": 1.0, "original_max_position_embeddings": 16384,
                                "rope_theta": 1000000.0, "rope_type": "yarn", "type": "yarn"}}
    vision = {"model_type": "pixtral", "hidden_size": 32, "intermediate_size": 64, "num_hidden_layers": 1,
              "num_attention_heads": 2, "head_dim": 16, "image_size": 1540, "patch_size": 14, "num_channels": 3,
              "rope_parameters": {"rope_theta": 10000.0, "rope_type": "default"}, "rope_theta": 10000.0}
    return Mistral3Config(text_config=text, vision_config=vision, image_token_index=10, spatial_merge_size=2,
                          multimodal_projector_bias=False, projector_hidden_act="gelu", vision_feature_layer=-1,
                          tie_word_embeddings=False)


class SchemaModel(nn.Module):
    """Backbone (optionally LoRA-wrapped) + head. `forward_micro` runs one micro-batch of encoded requests."""

    def __init__(self, backbone, head, taps, peft_model=None):
        super().__init__()
        self.backbone = backbone          # Mistral3ForConditionalGeneration (LoRA layers injected in place)
        self.peft_model = peft_model      # the PeftModel wrapper (for save_pretrained) or None
        self.head = head
        self.taps = list(taps)
        self._captured = {}
        self._capturing = False   # hooks record only during our own forward, not during checkpoint recompute
        self._hooks = []
        layers = decoder_layers(backbone)
        for t in self.taps:
            if t == "final":
                continue
            idx = int(t)
            if not 0 <= idx < len(layers):
                raise ValueError(f"tap layer {idx} out of range (0..{len(layers) - 1})")
            self._hooks.append(layers[idx].register_forward_hook(self._make_hook(idx)))

    def _make_hook(self, idx):
        def hook(module, inputs, output):
            if self._capturing:
                self._captured[idx] = output[0] if isinstance(output, (tuple, list)) else output
        return hook

    def device(self):
        return next(self.head.parameters()).device

    def backbone_states(self, batch):
        """Run the inner model on a padded batch; returns list of [B, T, H] tensors in tap order."""
        self._captured = {}
        dev = batch["input_ids"].device
        kwargs = {k: v for k, v in batch.items() if k in ("input_ids", "attention_mask", "position_ids",
                                                           "pixel_values", "image_sizes")}
        use_bf16 = dev.type == "cuda"
        self._capturing = True
        try:
            with torch.autocast(device_type=dev.type, dtype=torch.bfloat16, enabled=use_bf16):
                out = self.backbone.model(**kwargs, use_cache=False, return_dict=True)
        finally:
            self._capturing = False
        states = []
        for t in self.taps:
            states.append(out.last_hidden_state if t == "final" else self._captured[int(t)])
        self._captured = {}
        return states

    def forward_micro(self, encs, batch):
        """encs: list of encoded requests (render.Encoder.encode outputs, right-padded in `batch`).
        Returns list (per request) of lists (per question) of fp32 logit vectors (canonical option order)."""
        states = self.backbone_states(batch)
        outputs = []
        dev = states[0].device
        with torch.autocast(device_type=dev.type, enabled=False):
            for i, enc in enumerate(encs):
                n = enc["n_tokens"]
                taps = [s[i, :n] for s in states]   # bf16 on GPU; the head casts inside its checkpointed projection
                roles = torch.tensor(enc.get("_roles") or token_roles(enc), device=dev)
                outputs.append(self.head.forward_one(taps, enc, roles))
        return outputs

    def trainable_named(self):
        return [(n, p) for n, p in self.named_parameters() if p.requires_grad]


def pad_pixel_values(chunks):
    """Join per-row pixel tensors [n_i, C, H_i, W_i] into one [sum n_i, C, maxH, maxW] tensor, zero-padded at the
    bottom/right. Every row went through the processor on its own, so H/W differ between rows; this is exactly what
    PixtralImageProcessor._pad_for_batching does for a multi-image call, and the vision tower crops each image back
    to its `image_sizes` entry (patch_embeds[..., :h // patch, :w // patch]) so the padding never reaches the model."""
    h = max(int(c.shape[-2]) for c in chunks)
    w = max(int(c.shape[-1]) for c in chunks)
    return torch.cat([torch.nn.functional.pad(c, (0, w - int(c.shape[-1]), 0, h - int(c.shape[-2])))
                      for c in chunks], 0)


def make_batch(encs, pad_id, device):
    """Right-padded text batch (position ids restart per row); image micro-batches carry pixel values."""
    longest = max(e["n_tokens"] for e in encs)
    ids = torch.full((len(encs), longest), pad_id, dtype=torch.long)
    mask = torch.zeros((len(encs), longest), dtype=torch.long)
    for i, e in enumerate(encs):
        ids[i, :e["n_tokens"]] = torch.tensor(e["input_ids"])
        mask[i, :e["n_tokens"]] = 1
    pos = torch.arange(longest).unsqueeze(0).expand(len(encs), -1).contiguous()
    batch = {"input_ids": ids.to(device), "attention_mask": mask.to(device), "position_ids": pos.to(device)}
    pix = [e.get("_pixel") for e in encs]
    if any(p is not None for p in pix):
        if not all(p is not None for p in pix):
            raise ValueError("image micro-batches must be homogeneous")
        batch["pixel_values"] = pad_pixel_values([torch.as_tensor(p["pixel_values"]) for p in pix]).to(device)
        batch["image_sizes"] = torch.cat([torch.as_tensor(p["image_sizes"]) for p in pix], 0).to(device)
        if device.type == "cuda":
            batch["pixel_values"] = batch["pixel_values"].to(torch.bfloat16)
    return batch


def build_model(snapshot=None, tiny=None, lora_rank=64, lora_alpha=None, lora_dropout=0.0, taps=("final", 25),
                head_overrides=None, device="cpu", grad_checkpointing=True, init_adapter=None, attn="sdpa",
                seed=0):
    """Load (or create tiny) backbone, delete lm_head, freeze, add LoRA (rank 0 = none), attach the head.

    Returns (SchemaModel, info dict)."""
    import peft
    import transformers
    torch.manual_seed(seed)
    if tiny is not None:
        torch.manual_seed(1234)   # fixed: the tiny random "base" must be identical in training and serving
        cfg = tiny if not isinstance(tiny, dict) else tiny_config(**tiny)
        backbone = transformers.Mistral3ForConditionalGeneration._from_config(cfg, attn_implementation=attn)
        backbone = backbone.to(device)
    else:
        backbone = transformers.Mistral3ForConditionalGeneration.from_pretrained(
            str(snapshot), local_files_only=True, trust_remote_code=False, dtype=torch.bfloat16,
            attn_implementation=attn, device_map={"": device} if str(device) != "cpu" else None)
        if str(device) == "cpu":
            backbone = backbone.to(device)
    torch.manual_seed(seed)
    backbone.config.use_cache = False
    lm_head_bytes = backbone.lm_head.weight.numel() * backbone.lm_head.weight.element_size()
    del backbone.lm_head
    backbone.lm_head = nn.Identity()  # never called; keeps attribute access working
    for p in backbone.parameters():
        p.requires_grad_(False)
    if grad_checkpointing:
        backbone.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
        vt = getattr(backbone.model, "vision_tower", None)
        if vt is not None and hasattr(vt, "gradient_checkpointing_disable"):
            vt.gradient_checkpointing_disable()
    leaves = lora_leaves(backbone)
    peft_model = None
    if init_adapter:
        peft_model = peft.PeftModel.from_pretrained(backbone, str(init_adapter), is_trainable=True)
    elif lora_rank:
        alpha = lora_alpha if lora_alpha is not None else 2 * lora_rank
        peft_model = peft.get_peft_model(backbone, peft.LoraConfig(
            r=lora_rank, lora_alpha=alpha, lora_dropout=lora_dropout, target_modules=leaves,
            exclude_modules=VISION_EXCLUDE_REGEX, bias="none"))
    leaked = 0
    for name, p in backbone.named_parameters():
        if "vision_tower" in name or "multi_modal_projector" in name:
            if p.requires_grad:
                leaked += 1
            p.requires_grad_(False)
        elif p.requires_grad:
            p.data = p.data.float()  # LoRA weights fp32 (as the trainer)
    text_cfg = backbone.config.text_config
    hcfg = HeadConfig(hidden_size=text_cfg.hidden_size, n_taps=len(taps), taps=list(taps),
                      **(head_overrides or {}))
    head = JointSchemaHead(hcfg).to(device=device, dtype=torch.float32)
    model = SchemaModel(backbone, head, taps, peft_model)
    info = {"lora_leaves": leaves, "lora_rank": lora_rank,
            "lora_alpha": (lora_alpha if lora_alpha is not None else 2 * lora_rank) if lora_rank else None,
            "lora_params": sum(p.numel() for n, p in backbone.named_parameters() if p.requires_grad),
            "head_params": sum(p.numel() for p in head.parameters()), "lm_head_bytes_freed": lm_head_bytes,
            "vision_lora_leaked_frozen": leaked, "taps": list(taps), "head_config": hcfg.to_dict(),
            "layers": len(decoder_layers(backbone)), "hidden_size": text_cfg.hidden_size}
    return model, info


def save_head(head, folder, extra=None):
    from safetensors.torch import save_file
    folder = Path(folder)
    folder.mkdir(parents=True, exist_ok=True)
    save_file({k: v.detach().contiguous().cpu() for k, v in head.state_dict().items()},
              str(folder / "schema_head.safetensors"))
    cfg = {"head": head.cfg.to_dict(), **(extra or {})}
    tmp = folder / "schema_head_config.json.tmp"
    tmp.write_text(json.dumps(cfg, indent=2) + "\n")
    os.replace(tmp, folder / "schema_head_config.json")


def load_head_weights(head, folder):
    from safetensors.torch import load_file
    state = load_file(str(Path(folder) / "schema_head.safetensors"))
    missing, unexpected = head.load_state_dict(state, strict=True), None
    return missing