#!/usr/bin/env python3 """Step 3: train a LoRA that carries capability C, with the base FROZEN. A small synthetic dataset (a few hundred request -> directive pairs) trains a PEFT LoRA adapter. The base weights are frozen: only the adapter learns. To prove the base is untouched, every base weight shard in the Hugging Face cache is SHA-256 hashed before and after training and the digests are shown to be byte-identical. The trained adapter is written to a gitignored output dir. """ import hashlib import json import sys import time from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from demo.common import ( ADAPTER_DIR, BASE_MODEL_ID, STATE_DIR, base_weight_files, build_prompt, load_tokenizer, ) from demo.dataset import build_dataset def hash_files(paths): digests = {} for path in paths: h = hashlib.sha256() h.update(path.read_bytes()) digests[path.name] = h.hexdigest() return digests def main() -> int: import torch from peft import LoraConfig, get_peft_model from transformers import AutoModelForCausalLM print("=== Step 3: train LoRA for capability C (base frozen) ===\n") torch.manual_seed(0) train, _held_out = build_dataset() print(f" training examples: {len(train)}") tokenizer = load_tokenizer() # --- hash the base weights BEFORE training --- base_before = hash_files(base_weight_files()) print(f" base weight shards: {list(base_before)}") model = AutoModelForCausalLM.from_pretrained(BASE_MODEL_ID, dtype=torch.float32) model.to("cuda") lora = LoraConfig( r=16, lora_alpha=32, lora_dropout=0.05, target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"], task_type="CAUSAL_LM", ) model = get_peft_model(model, lora) trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) total = sum(p.numel() for p in model.parameters()) frozen_base = all( (not p.requires_grad) for n, p in model.named_parameters() if "lora_" not in n ) print(f" trainable params: {trainable:,} of {total:,} " f"({100 * trainable / total:.3f}% -- the LoRA only)") print(f" every non-LoRA (base) parameter has requires_grad=False: {frozen_base}\n") # --- build masked training tensors --- def encode(example): full_messages = [ {"role": "system", "content": _system()}, {"role": "user", "content": example.request}, {"role": "assistant", "content": example.directive.render()}, ] full_text = tokenizer.apply_chat_template( full_messages, tokenize=False, add_generation_prompt=False ) full_ids = tokenizer(full_text, add_special_tokens=False)["input_ids"] prompt_text = build_prompt(tokenizer, example.request) prompt_ids = tokenizer(prompt_text, add_special_tokens=False)["input_ids"] labels = list(full_ids) for i in range(min(len(prompt_ids), len(labels))): labels[i] = -100 # mask the prompt; train only on the directive return full_ids, labels encoded = [encode(e) for e in train] def collate(batch): max_len = max(len(ids) for ids, _ in batch) pad_id = tokenizer.pad_token_id input_ids, attn, label_ids = [], [], [] for ids, labels in batch: pad = max_len - len(ids) input_ids.append(ids + [pad_id] * pad) attn.append([1] * len(ids) + [0] * pad) label_ids.append(labels + [-100] * pad) return ( torch.tensor(input_ids), torch.tensor(attn), torch.tensor(label_ids), ) epochs = 3 batch_size = 8 lr = 2e-4 optimizer = torch.optim.AdamW( [p for p in model.parameters() if p.requires_grad], lr=lr ) model.train() start = time.time() step = 0 losses = [] for epoch in range(epochs): order = torch.randperm(len(encoded)) for b in range(0, len(encoded), batch_size): idx = order[b : b + batch_size].tolist() input_ids, attn, label_ids = collate([encoded[i] for i in idx]) input_ids, attn, label_ids = ( input_ids.cuda(), attn.cuda(), label_ids.cuda() ) out = model(input_ids=input_ids, attention_mask=attn, labels=label_ids) out.loss.backward() optimizer.step() optimizer.zero_grad() step += 1 losses.append(out.loss.item()) if step % 20 == 0 or step == 1: print(f" epoch {epoch + 1}/{epochs} step {step:4d} " f"loss {out.loss.item():.4f}") train_seconds = time.time() - start print(f"\n training complete: {step} steps in {train_seconds:.1f}s " f"(final loss {losses[-1]:.4f})") # --- save the adapter (gitignored) --- ADAPTER_DIR.mkdir(parents=True, exist_ok=True) model.save_pretrained(str(ADAPTER_DIR)) adapter_file = ADAPTER_DIR / "adapter_model.safetensors" adapter_size = adapter_file.stat().st_size if adapter_file.exists() else 0 print(f" adapter saved to: {ADAPTER_DIR} ({adapter_size / 1e6:.2f} MB, gitignored)") # --- hash the base weights AFTER training --- base_after = hash_files(base_weight_files()) identical = base_before == base_after print("\n base-frozen proof (SHA-256 of base weight shards):") for name in base_before: same = base_before[name] == base_after.get(name) print(f" {name}: {base_before[name][:16]}... " f"{'UNCHANGED' if same else 'CHANGED'}") print(f" base weights byte-identical before and after training: {identical}") STATE_DIR.mkdir(parents=True, exist_ok=True) (STATE_DIR / "train_report.json").write_text(json.dumps({ "steps": step, "epochs": epochs, "batch_size": batch_size, "lr": lr, "train_seconds": round(train_seconds, 1), "final_loss": round(losses[-1], 4), "trainable_params": trainable, "total_params": total, "adapter_bytes": adapter_size, "base_frozen": identical and frozen_base, }, indent=2)) (STATE_DIR / "base_hashes.json").write_text(json.dumps({ "before": base_before, "after": base_after, "identical": identical }, indent=2)) if not (identical and frozen_base): print("\nRESULT: FAIL - base weights changed; training was not frozen.") return 1 print("\nRESULT: PASS - LoRA trained; base weights provably frozen.") return 0 def _system(): from demo.common import SYSTEM_PROMPT return SYSTEM_PROMPT if __name__ == "__main__": sys.exit(main())