File size: 13,869 Bytes
4be6a52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
"""Optional native Clef training primitives; imported only by ML workflows."""

from __future__ import annotations

import hashlib
import json
import math
import re
from pathlib import Path
from typing import Any, Literal

import torch
from safetensors.torch import load_file, save_file
from torch import nn
from torch.nn import functional as functional

from stackcraft.clef import ENCODING_VERSION, MODEL_ID, MODEL_REVISION, SOURCE_SHA256

FORMAT_VERSION = 1
HEAD_TYPE = "native-joint-schema-fp32-gathered-rows-v1"
_TARGET = re.compile(
    r"^model\.language_model\.layers\.\d+\."
    r"(?:self_attn|linear_attn|mlp)\."
    r"(?:q_proj|k_proj|v_proj|o_proj|in_proj_qkv|in_proj_z|in_proj_b|in_proj_a|"
    r"out_proj|gate_proj|up_proj|down_proj)$"
)


class GatheredFloat32Embedding:
    """The pinned head only indexes this object; never materialize all rows in FP32."""

    def __init__(self, weight: torch.Tensor) -> None:
        self.weight = weight

    def __getitem__(self, indices: torch.Tensor) -> torch.Tensor:
        # Both indexing and casting retain autograd edges when the input requires it.
        return self.weight[indices].float()


class FP32DecisionHead(nn.Module):
    """Keep the upstream head unchanged while adapting its floating inputs."""

    def __init__(self, native_head: nn.Module) -> None:
        super().__init__()
        self.native_head = native_head.float()

    def forward(
        self,
        hidden_states: torch.Tensor,
        input_ids: torch.Tensor,
        attention_mask: torch.Tensor,
        records: list[Any],
        output_embedding_weight: torch.Tensor,
    ) -> list[list[torch.Tensor]]:
        # Disable surrounding autocast so trainable head and its operations stay FP32.
        with torch.autocast(device_type=hidden_states.device.type, enabled=False):
            return self.native_head(
                hidden_states.float(),
                input_ids,
                attention_mask,
                records,
                GatheredFloat32Embedding(output_embedding_weight),
            )


def lora_target_modules(backbone: nn.Module) -> list[str]:
    """Return full text-layer names; suffix-only matching can accidentally train vision."""
    targets = [
        name
        for name, module in backbone.named_modules()
        if isinstance(module, nn.Linear) and _TARGET.fullmatch(name)
    ]
    if not targets:
        raise ValueError("no supported Qwen3.5 text-layer LoRA targets found")
    return sorted(targets)


def prepare_trainable(model: Any, mode: Literal["head", "lora"] = "head", rank: int = 4) -> Any:
    """Prepare a pinned, already admitted/loaded native ClefModel in place.

    This function loads no model or weights. The caller must supply the pinned
    native model, e.g. ClefPlayer.from_pretrained(...).model after memory admission.
    Backbone dtype is preserved; the intended release loader supplies BF16.
    """
    if mode not in ("head", "lora"):
        raise ValueError("training mode must be head or lora")
    if type(rank) is not int or rank < 1:
        raise ValueError("LoRA rank must be a positive integer")
    if isinstance(model.head, FP32DecisionHead) or hasattr(model, "_stackcraft_training"):
        raise ValueError("model is already prepared for Stackcraft training")
    if hasattr(model.language_model, "peft_config"):
        raise ValueError("expected the unchanged backbone, not an existing PEFT model")
    model.language_model.requires_grad_(False)
    targets: list[str] = []
    if mode == "lora":
        from peft import LoraConfig, get_peft_model

        targets = lora_target_modules(model.language_model)
        model.language_model = get_peft_model(
            model.language_model,
            LoraConfig(
                r=rank,
                lora_alpha=2 * rank,
                lora_dropout=0.0,
                target_modules=targets,
                bias="none",
            ),
        )
        if hasattr(model.language_model, "gradient_checkpointing_enable"):
            model.language_model.gradient_checkpointing_enable(
                gradient_checkpointing_kwargs={"use_reentrant": False}
            )
    model.head = FP32DecisionHead(model.head)
    model.head.requires_grad_(True)
    model.train()
    if mode == "head":
        model.language_model.eval()
    model._stackcraft_training = {
        "format_version": FORMAT_VERSION,
        "base_model": MODEL_ID,
        "base_revision": MODEL_REVISION,
        "native_source_sha256": SOURCE_SHA256,
        "encoding_version": ENCODING_VERSION,
        "head_type": HEAD_TYPE,
        "mode": mode,
        "lora": (
            {"rank": rank, "alpha": 2 * rank, "dropout": 0.0, "target_modules": targets}
            if mode == "lora"
            else None
        ),
        "loss": {"label_smoothing": 0.05, "brier_weight": 0.1, "brier_reduction": "sum_options"},
    }
    return model


def decision_loss(
    logits: torch.Tensor,
    encoded_record: Any,
    target_action_id: str,
    *,
    label_smoothing: float = 0.05,
    brier_weight: float = 0.1,
) -> torch.Tensor:
    """One native choice: smoothed cross entropy plus multiclass Brier sum."""
    if len(encoded_record.questions) != 1:
        raise ValueError("training records must contain exactly one choice question")
    question = encoded_record.questions[0]
    if question.question_type != 1:
        raise ValueError("training question must be native choice type 1")
    ids = question.option_ids
    if tuple(ids) != tuple(sorted(set(ids))):
        raise ValueError("encoded option IDs must be unique and lexicographically sorted")
    if target_action_id not in ids:
        raise ValueError("target action is missing from native encoded option IDs")
    if logits.ndim != 1 or logits.numel() != len(ids):
        raise ValueError("logits shape does not match the encoded choices")
    if (
        not math.isfinite(label_smoothing)
        or not 0 <= label_smoothing <= 1
        or not math.isfinite(brier_weight)
        or brier_weight < 0
    ):
        raise ValueError("invalid label smoothing or Brier weight")
    if not torch.isfinite(logits).all():
        raise ValueError("decision logits contain nonfinite values")
    values = logits.float().unsqueeze(0)
    target = torch.tensor([ids.index(target_action_id)], device=values.device)
    cross_entropy = functional.cross_entropy(values, target, label_smoothing=label_smoothing)
    one_hot = functional.one_hot(target, num_classes=len(ids)).float()
    brier = (values.softmax(-1) - one_hot).square().sum(-1).mean()
    return cross_entropy + brier_weight * brier


def parameter_hashes(
    model: nn.Module, *, trainable: bool, chunk_elements: int = 1_048_576
) -> dict[str, str]:
    """Hash selected parameters exactly, moving only bounded chunks to CPU.

    Full frozen-backbone hashing is intentionally an explicit before/after audit,
    not a training-step operation. Dtype and shape are included in every digest.
    """
    if type(chunk_elements) is not int or chunk_elements < 1:
        raise ValueError("chunk_elements must be a positive integer")
    results = {}
    for name, parameter in model.named_parameters():
        if parameter.requires_grad != trainable:
            continue
        digest = hashlib.sha256()
        digest.update(f"{parameter.dtype}:{tuple(parameter.shape)}:".encode())
        flattened = parameter.detach().reshape(-1)
        for start in range(0, flattened.numel(), chunk_elements):
            chunk = flattened[start : start + chunk_elements].to("cpu").contiguous()
            digest.update(chunk.view(torch.uint8).numpy().tobytes())
        results[name] = digest.hexdigest()
    return results


def save_checkpoint(
    model: Any, path: str | Path, *, extra_metadata: dict[str, Any] | None = None
) -> dict[str, Any]:
    """Save head and optional LoRA separately, never the frozen multi-GB backbone."""
    if not isinstance(model.head, FP32DecisionHead) or not hasattr(model, "_stackcraft_training"):
        raise ValueError("model must be prepared before saving a training checkpoint")
    destination = Path(path)
    destination.mkdir(parents=True, exist_ok=False)
    metadata = dict(model._stackcraft_training)
    metadata["extra"] = extra_metadata or {}
    # Store native head keys, not wrapper-specific state_dict prefixes.
    head_state = {
        name: tensor.detach().cpu().contiguous()
        for name, tensor in model.head.native_head.state_dict().items()
    }
    save_file(head_state, destination / "joint_head.safetensors")
    metadata["head_shapes"] = {name: list(tensor.shape) for name, tensor in head_state.items()}
    if metadata["mode"] == "lora":
        model.language_model.save_pretrained(destination / "adapter", safe_serialization=True)
        adapter_path = destination / "adapter" / "adapter_config.json"
        adapter_config = json.loads(adapter_path.read_text())
        adapter_config.update(
            base_model_name_or_path=MODEL_ID,
            revision=MODEL_REVISION,
            target_modules=metadata["lora"]["target_modules"],
        )
        adapter_path.write_text(json.dumps(adapter_config, indent=2, sort_keys=True) + "\n")
    (destination / "training_config.json").write_text(
        json.dumps(metadata, indent=2, sort_keys=True, allow_nan=False) + "\n"
    )
    return metadata


def load_checkpoint(model: Any, path: str | Path, *, trainable: bool = False) -> Any:
    """Restore onto an unchanged pinned native base; reject incompatible metadata."""
    source = Path(path)
    metadata = json.loads((source / "training_config.json").read_text())
    expected = {
        "format_version": FORMAT_VERSION,
        "base_model": MODEL_ID,
        "base_revision": MODEL_REVISION,
        "native_source_sha256": SOURCE_SHA256,
        "encoding_version": ENCODING_VERSION,
        "head_type": HEAD_TYPE,
    }
    for key, value in expected.items():
        if type(metadata.get(key)) is not type(value) or metadata[key] != value:
            raise ValueError(f"checkpoint {key} is incompatible with this pinned native adapter")
    mode = metadata.get("mode")
    if mode not in ("head", "lora"):
        raise ValueError("checkpoint training mode is invalid")
    if isinstance(model.head, FP32DecisionHead) or hasattr(model.language_model, "peft_config"):
        raise ValueError("checkpoint must load onto an unchanged native base")
    head_state = load_file(source / "joint_head.safetensors", device="cpu")
    shapes = {name: list(tensor.shape) for name, tensor in head_state.items()}
    base_shapes = {name: list(tensor.shape) for name, tensor in model.head.state_dict().items()}
    if shapes != metadata.get("head_shapes") or shapes != base_shapes:
        raise ValueError("checkpoint head structure differs from metadata or native model")
    if any(tensor.dtype != torch.float32 for tensor in head_state.values()):
        raise ValueError("checkpoint head tensors must be FP32")
    model.language_model.requires_grad_(False)
    if mode == "lora":
        from peft import LoraConfig, PeftModel
        from peft.tuners.tuners_utils import check_target_module_exists

        lora = metadata.get("lora")
        if not isinstance(lora, dict) or lora.get("target_modules") != lora_target_modules(
            model.language_model
        ):
            raise ValueError("checkpoint LoRA targets differ from the native text backbone")
        if (
            type(lora.get("rank")) is not int
            or lora["rank"] < 1
            or lora.get("alpha") != 2 * lora["rank"]
            or lora.get("dropout") != 0.0
        ):
            raise ValueError(
                "checkpoint LoRA rank, alpha or dropout violates the training contract"
            )
        config = json.loads((source / "adapter" / "adapter_config.json").read_text())
        # PEFT 0.21.2 minimizes >=20 explicit module names to equivalent suffixes.
        # Compare their meaning on this exact unchanged backbone, not list spelling.
        # Enumerating ALL modules ensures an accidental vision/MTP/lm_head match
        # makes the sets unequal and is rejected before installing the adapter.
        saved_config = LoraConfig.from_pretrained(str(source / "adapter"))
        resolved_targets = sorted(
            name
            for name, _ in model.language_model.named_modules()
            if check_target_module_exists(saved_config, name)
        )
        if (
            config.get("r") != lora.get("rank")
            or config.get("lora_alpha") != lora.get("alpha")
            or config.get("lora_dropout") != lora.get("dropout")
            or resolved_targets != lora["target_modules"]
            or config.get("bias") != "none"
            or config.get("modules_to_save") is not None
            or config.get("target_parameters") is not None
        ):
            raise ValueError("saved adapter configuration differs from checkpoint metadata")
        model.language_model = PeftModel.from_pretrained(
            model.language_model, source / "adapter", is_trainable=trainable
        )
        if trainable and hasattr(model.language_model, "gradient_checkpointing_enable"):
            model.language_model.gradient_checkpointing_enable(
                gradient_checkpointing_kwargs={"use_reentrant": False}
            )
    elif metadata.get("lora") is not None:
        raise ValueError("head-only checkpoint must not contain LoRA configuration")
    model.head = FP32DecisionHead(model.head)
    model.head.native_head.load_state_dict(head_state, strict=True)
    model.head.requires_grad_(trainable)
    model._stackcraft_training = {
        key: value for key, value in metadata.items() if key not in ("extra", "head_shapes")
    }
    model.train(trainable)
    if mode == "head":
        model.language_model.eval()
    return model