stackcraft-clef-flash-lora / code /tests /test_training.py
nima1's picture
Publish verified checkpoint and losslessly compressed study evidence
4be6a52 verified
Raw History Blame Contribute Delete
15.2 kB
"""Tiny native-head training checks, not proof that the 9B GPU probe fits."""
import copy
import json
from pathlib import Path
from types import SimpleNamespace
import pytest
torch = pytest.importorskip("torch", reason="training tests require the optional ML extra")
pytest.importorskip("peft", reason="training tests require the optional ML extra")
from stackcraft.clef import MODEL_ID, MODEL_REVISION, import_pinned_source # noqa: E402
from stackcraft.training import ( # noqa: E402
FP32DecisionHead,
GatheredFloat32Embedding,
decision_loss,
load_checkpoint,
lora_target_modules,
parameter_hashes,
prepare_trainable,
save_checkpoint,
)
@pytest.fixture(scope="module")
def native():
from huggingface_hub import hf_hub_download
try:
path = hf_hub_download(
MODEL_ID, "joint_schema_model.py", revision=MODEL_REVISION, local_files_only=True
)
except OSError:
pytest.skip("pinned native source must be cached for actual-head CPU tests")
return import_pinned_source(Path(path), trust_pinned_code=True)
def record(native):
return native.EncodedRecord(
input_ids=tuple(range(16)),
questions=(
native.EncodedQuestion(
question_id="placement",
question_type=1,
question_span=(1, 3),
option_spans=((4, 6), (7, 9), (10, 12)),
option_ids=("r0x0", "r0x1", "r1x0"),
),
),
record_id="tiny-test",
)
class TinyText(torch.nn.Module):
def __init__(self):
super().__init__()
self.embed_tokens = torch.nn.Embedding(32, 8)
layer = torch.nn.Module()
layer.self_attn = torch.nn.Module()
layer.self_attn.q_proj = torch.nn.Linear(8, 8, bias=False)
layer.linear_attn = torch.nn.Module()
layer.linear_attn.in_proj_qkv = torch.nn.Linear(8, 8, bias=False)
layer.linear_attn.in_proj_z = torch.nn.Linear(8, 8, bias=False)
layer.mlp = torch.nn.Module()
layer.mlp.down_proj = torch.nn.Linear(8, 8, bias=False)
self.layers = torch.nn.ModuleList([layer])
def forward(self, input_ids, **kwargs):
hidden = self.embed_tokens(input_ids)
for layer in self.layers:
hidden = hidden + layer.mlp.down_proj(
(
layer.self_attn.q_proj(hidden)
+ layer.linear_attn.in_proj_qkv(hidden)
+ layer.linear_attn.in_proj_z(hidden)
).tanh()
)
return SimpleNamespace(last_hidden_state=hidden)
class TinyBackbone(torch.nn.Module):
def __init__(self):
super().__init__()
self.config = {"model_type": "custom", "_name_or_path": "tiny-offline-test"}
self.model = torch.nn.Module()
self.model.language_model = TinyText()
self.model.visual = torch.nn.Module()
self.model.visual.q_proj = torch.nn.Linear(8, 8)
self.mtp = torch.nn.Module()
self.mtp.q_proj = torch.nn.Linear(8, 8)
self.lm_head = torch.nn.Linear(8, 32, bias=False)
def get_output_embeddings(self):
return self.lm_head
def tiny_model(native):
torch.manual_seed(42)
head = native.JointSchemaHead(
hidden_size=8, width=8, routing_layers=1, layers=1, heads=2, feedforward=16
)
return native.ClefModel(TinyBackbone().to(torch.bfloat16), head.to(torch.bfloat16))
def test_gather_cast_matches_full_cast_and_preserves_selected_gradients():
weight = torch.randn(32, 8, dtype=torch.bfloat16, requires_grad=True)
selected = torch.tensor([2, 7, 2])
output = GatheredFloat32Embedding(weight)[selected]
assert torch.equal(output, weight.float()[selected])
output.sum().backward()
assert torch.all(weight.grad[2] == 2)
assert torch.all(weight.grad[7] == 1)
assert torch.count_nonzero(weight.grad).item() == 16
def test_actual_native_head_float32_wrapper_matches_reference_and_gradients(native):
model = tiny_model(native)
head = model.head.float().eval()
wrapper = FP32DecisionHead(copy.deepcopy(head)).eval()
encoded = record(native)
batch = native.collate_records([encoded], 0, torch.device("cpu"))
hidden = torch.randn(1, 16, 8, dtype=torch.bfloat16, requires_grad=True)
reference_hidden = hidden.detach().clone().requires_grad_(True)
weight = torch.randn(32, 8, dtype=torch.bfloat16, requires_grad=True)
reference_weight = weight.detach().clone().requires_grad_(True)
wrapped = wrapper(hidden, batch["input_ids"], batch["attention_mask"], [encoded], weight)[0][0]
reference = head(
reference_hidden.float(),
batch["input_ids"],
batch["attention_mask"],
[encoded],
reference_weight.float(),
)[0][0]
assert wrapped.dtype == torch.float32
torch.testing.assert_close(wrapped, reference, rtol=0, atol=0)
wrapped.square().sum().backward()
reference.square().sum().backward()
torch.testing.assert_close(hidden.grad, reference_hidden.grad, rtol=0, atol=0)
torch.testing.assert_close(weight.grad, reference_weight.grad, rtol=0, atol=0)
assert torch.count_nonzero(hidden.grad) > 0
for (_, actual), (_, expected) in zip(
wrapper.native_head.named_parameters(), head.named_parameters(), strict=True
):
torch.testing.assert_close(actual.grad, expected.grad, rtol=0, atol=0)
def test_lora_targets_use_full_names_and_exclude_vision_output_and_mtp(native):
targets = lora_target_modules(tiny_model(native).language_model)
assert targets == [
"model.language_model.layers.0.linear_attn.in_proj_qkv",
"model.language_model.layers.0.linear_attn.in_proj_z",
"model.language_model.layers.0.mlp.down_proj",
"model.language_model.layers.0.self_attn.q_proj",
]
@pytest.mark.parametrize("mode", ["head", "lora"])
def test_real_native_forward_gradients_changed_intended_bytes_and_reload(native, tmp_path, mode):
model = prepare_trainable(tiny_model(native), mode=mode)
encoded = record(native)
batch = native.collate_records([encoded], 0, torch.device("cpu"))
assert isinstance(model.head, FP32DecisionHead)
assert all(parameter.dtype == torch.float32 for parameter in model.head.parameters())
assert all(
parameter.dtype == torch.bfloat16
for name, parameter in model.language_model.named_parameters()
if "lora_" not in name
)
before = parameter_hashes(model, trainable=True, chunk_elements=7)
frozen = parameter_hashes(model, trainable=False, chunk_elements=7)
optimizer = torch.optim.AdamW((p for p in model.parameters() if p.requires_grad), lr=0.002)
for _ in range(2):
optimizer.zero_grad(set_to_none=True)
logits = model(batch)[0][0]
loss = decision_loss(logits, encoded, "r0x1")
assert torch.isfinite(loss)
loss.backward()
grads = [p.grad for p in model.head.parameters() if p.grad is not None]
assert all(torch.isfinite(grad).all() for grad in grads)
assert any(torch.count_nonzero(grad) > 0 for grad in grads)
assert all(p.grad is None for p in model.parameters() if not p.requires_grad)
optimizer.step()
after = parameter_hashes(model, trainable=True, chunk_elements=7)
assert any(before[name] != digest for name, digest in after.items() if name.startswith("head."))
if mode == "lora":
assert any(before[name] != digest for name, digest in after.items() if "lora_" in name)
assert parameter_hashes(model, trainable=False, chunk_elements=7) == frozen
model.eval()
with torch.no_grad():
expected = model(batch)[0][0].softmax(-1)
destination = tmp_path / mode
metadata = save_checkpoint(model, destination, extra_metadata={"scope": "tiny-cpu-plumbing"})
assert metadata["mode"] == mode
assert (destination / "joint_head.safetensors").exists()
assert (destination / "adapter").exists() == (mode == "lora")
if mode == "lora":
adapter_config = json.loads((destination / "adapter" / "adapter_config.json").read_text())
assert adapter_config["base_model_name_or_path"] == MODEL_ID
assert adapter_config["revision"] == MODEL_REVISION
assert adapter_config["target_modules"] == metadata["lora"]["target_modules"]
restored = load_checkpoint(tiny_model(native), destination)
assert not any(parameter.requires_grad for parameter in restored.parameters())
with torch.no_grad():
actual = restored(batch)[0][0].softmax(-1)
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
with pytest.raises(FileExistsError):
save_checkpoint(model, destination)
def test_loss_uses_sorted_choice_index_and_declared_formula(native):
encoded = record(native)
logits = torch.tensor([0.2, -0.4, 1.2], requires_grad=True)
actual = decision_loss(logits, encoded, "r0x1")
ce = torch.nn.functional.cross_entropy(
logits.unsqueeze(0), torch.tensor([1]), label_smoothing=0.05
)
brier = ((logits.softmax(-1) - torch.tensor([0.0, 1.0, 0.0])) ** 2).sum()
torch.testing.assert_close(actual, ce + 0.1 * brier)
actual.backward()
assert torch.count_nonzero(logits.grad) > 0
with pytest.raises(ValueError, match="missing"):
decision_loss(logits, encoded, "not-an-action")
with pytest.raises(ValueError, match="nonfinite"):
decision_loss(torch.tensor([0.0, float("nan"), 1.0]), encoded, "r0x1")
@pytest.mark.parametrize(
"key,value",
[
("format_version", True),
("format_version", 2),
("base_model", "other/model"),
("base_revision", "wrong"),
("native_source_sha256", "wrong"),
("encoding_version", "wrong"),
("head_type", "wrong"),
],
)
def test_checkpoint_rejects_wrong_contract(native, tmp_path, key, value):
model = prepare_trainable(tiny_model(native))
checkpoint = tmp_path / "checkpoint"
metadata = save_checkpoint(model, checkpoint)
metadata[key] = value
(checkpoint / "training_config.json").write_text(json.dumps(metadata))
with pytest.raises(ValueError, match=key):
load_checkpoint(tiny_model(native), checkpoint)
def test_parameter_hashes_are_chunk_independent_and_track_dtype_and_values():
model = torch.nn.Linear(4, 3).to(torch.bfloat16)
first = parameter_hashes(model, trainable=True, chunk_elements=1)
assert first == parameter_hashes(model, trainable=True, chunk_elements=10)
with torch.no_grad():
model.weight[0, 0] += 1
second = parameter_hashes(model, trainable=True)
assert first["weight"] != second["weight"]
assert first["bias"] == second["bias"]
assert first["bias"] != parameter_hashes(model.float(), trainable=True)["bias"]
def test_actual_tiny_qwen_hybrid_checkpointing_keeps_lora_gradients(native):
from transformers import Qwen3_5Config, Qwen3_5ForConditionalGeneration
config = Qwen3_5Config(
text_config={
"vocab_size": 32,
"hidden_size": 8,
"intermediate_size": 16,
"num_hidden_layers": 2,
"num_attention_heads": 1,
"num_key_value_heads": 1,
"head_dim": 8,
"layer_types": ["linear_attention", "full_attention"],
"linear_key_head_dim": 8,
"linear_value_head_dim": 8,
"linear_num_key_heads": 1,
"linear_num_value_heads": 1,
"max_position_embeddings": 64,
"use_cache": False,
"rope_parameters": {
"rope_type": "default",
"rope_theta": 10000,
"partial_rotary_factor": 1.0,
"mrope_section": [1, 1, 2],
},
},
vision_config={
"depth": 1,
"hidden_size": 8,
"intermediate_size": 16,
"num_heads": 1,
"patch_size": 2,
"spatial_merge_size": 1,
"temporal_patch_size": 1,
"out_hidden_size": 8,
"num_position_embeddings": 16,
},
)
backbone = Qwen3_5ForConditionalGeneration(config).to(torch.bfloat16)
targets = lora_target_modules(backbone)
assert len(targets) == 15 # 5 linear-attention + 4 full-attention + 6 MLP projections.
assert sum("linear_attn" in name for name in targets) == 5
assert all(name.startswith("model.language_model.layers.") for name in targets)
head = tiny_model(native).head
model = prepare_trainable(native.ClefModel(backbone, head), mode="lora")
assert model.language_model.get_base_model().is_gradient_checkpointing
encoded = record(native)
batch = native.collate_records([encoded], 0, torch.device("cpu"))
loss = decision_loss(model(batch)[0][0], encoded, "r0x1")
assert torch.isfinite(loss)
loss.backward()
for part in ("linear_attn", "self_attn", "mlp"):
assert any(
parameter.grad is not None and torch.count_nonzero(parameter.grad) > 0
for name, parameter in model.named_parameters()
if "lora_" in name and part in name
)
def test_peft_minimized_suffixes_reload_but_extra_targets_fail(native, tmp_path):
def larger_tiny_model():
model = tiny_model(native)
text = model.language_model.model.language_model
text.layers = torch.nn.ModuleList([copy.deepcopy(text.layers[0]) for _ in range(32)])
return model
model = prepare_trainable(larger_tiny_model(), mode="lora")
compressed = sorted(model.language_model.peft_config["default"].target_modules)
intended = model._stackcraft_training["lora"]["target_modules"]
assert len(intended) == 128
assert len(compressed) < len(intended) # Exercise actual PEFT >=20 optimization.
checkpoint = tmp_path / "compressed"
save_checkpoint(model, checkpoint)
config_path = checkpoint / "adapter" / "adapter_config.json"
config = json.loads(config_path.read_text())
assert config["target_modules"] == intended # New saves are explicit and canonical.
config["target_modules"] = compressed # Recreate the immutable older probe format.
config["base_model_name_or_path"] = "/old/local/cache/snapshot"
config["revision"] = None
config_path.write_text(json.dumps(config))
restored = load_checkpoint(larger_tiny_model(), checkpoint)
original_adapters = {
name: parameter.detach() for name, parameter in model.named_parameters() if "lora_" in name
}
restored_adapters = {
name: parameter.detach()
for name, parameter in restored.named_parameters()
if "lora_" in name
}
assert original_adapters.keys() == restored_adapters.keys()
for name, expected in original_adapters.items():
torch.testing.assert_close(restored_adapters[name], expected, rtol=0, atol=0)
# Force an otherwise-valid suffix to also select the output layer.
config["target_modules"] = compressed + ["lm_head"]
config_path.write_text(json.dumps(config))
with pytest.raises(ValueError, match="adapter configuration"):
load_checkpoint(larger_tiny_model(), checkpoint)