#%%writefile train_toolace_lora_ultra.py # ============================================================================== # JiRack Ultra ToolACE LoRA SFT + merge (single script, all four sizes) # COPYRIGHT (c) 2026 Konstantin Vladimirovich Grabko. # # Adapted from the JiRackPrecision_8b ToolACE LoRA script. One file covers # Ultra 1B / 7B / 14B / 32B -- set SIZE below, everything else follows from # the SIZES table. # # What it does (unchanged from the 8B original): # 1. Loads JiRackTransformer + your .pt checkpoint. # 2. Freezes everything; injects LoRA (A/B low-rank pairs) into every # nn.Linear except the LM head (out_features == vocab_size). # 3. ALSO unfreezes the embedding rows of the JiRack special tokens -- # those rows are untrained padded slots right now; the model can't emit # <|tool_call_start|> etc. until they're trained. A gradient hook zeroes # grads for all other rows, so the base vocab embeddings stay untouched. # 4. Trains with assistant-only loss masking, bf16 autocast, grad accum. # 5. Saves the LoRA adapter alone + OPTIONAL merged checkpoint whose # state_dict keys match the input .pt exactly. # # ============================ ULTRA-SPECIFIC CHANGES ========================== # [U-A] BitLinear IS an nn.Linear subclass in every Ultra file, so isinstance() # picks it up and LoRA wraps it -- which is what we want. But it also # means the wrapped base runs BitLinear.forward, i.e. the quantization # math, on every call. At LAMBDA=0.0 the result is mathematically # identical to plain F.linear (lam=0 => w_effective=w, x_effective=x), # but the BitLinear fast path only triggers in EVAL mode, so during # training you pay for the quant math with no effect. Tolerable; if you # want it gone, set LAMBDA=0.0 and patch BitLinear's fast-path condition # to also fire while training. # [U-B] FREEZE_8BIT IS DISABLED for Ultra. The 8B original swapped frozen # nn.Linear for bnb.nn.Linear8bitLt -- on Ultra that would REPLACE your # BitLinear modules with plain bnb layers, destroying the lambda_ buffers # and the ternary path, and the merged state_dict keys would no longer # match your checkpoint. Also, the original's merge_into_base() called # bnb.functional.dequantize_4bit on an 8-bit layer, which is the wrong # function anyway. If you need the memory, use adafactor + shorter # MAX_LEN, or shard -- not this. # [U-C] Tokenizer defaults to CMSManhattan/JiRackPrecisionTokenizer (the # published JiRack tokenizer), with a hard assert that it FITS the padded # matrix. Never resize: 151,779 < 151,936 (1B) and < 152,064 (7/14/32B), # so a resize would SHRINK and corrupt the embedding matrix. # [U-D] Per-size memory defaults in the SIZES table (MAX_LEN, GRAD_ACCUM). # ============================================================================== import json import math import os import random import sys import time import torch import torch.nn as nn from transformers import AutoTokenizer from transformers.optimization import Adafactor sys.path.append(os.getcwd()) # ========================= PICK YOUR SIZE ========================= SIZE = "1b" # "1b" | "7b" | "14b" | "32b" # ================================================================== # NOTE the module names -- they are NOT uniform in your repo: # 1B -> JiRackTernaryUltra_1b.py # 7B -> JiRackTernaryUltra7b.py <-- no underscore before "7b"! # 14B -> JiRackTernaryUltra_14b.py # 32B -> JiRackTernaryUltra_32b.py # If you rename any of them, fix the "module" field below. SIZES = { "1b": { "module": "JiRackTernaryUltra_1b", "vocab": 151936, "model_path": "/mnt/nfs_clientshare/JiRackUltra_1b/model.pt", "adapter": "/mnt/nfs_clientshare/JiRackUltra_1b/toolace_lora_adapter.pt", "merged": "/mnt/nfs_clientshare/JiRackUltra_1b/ultra1b_toolace.pt", "max_len": 2048, "grad_accum": 8, }, "7b": { "module": "JiRackTernaryUltra7b", "vocab": 152064, "model_path": "/mnt/nfs_clientshare/JiRackUltra_7b/model.pt", "adapter": "/mnt/nfs_clientshare/JiRackUltra_7b/toolace_lora_adapter.pt", "merged": "/mnt/nfs_clientshare/JiRackUltra_7b/ultra7b_toolace.pt", "max_len": 2048, "grad_accum": 16, }, "14b": { "module": "JiRackTernaryUltra_14b", "vocab": 152064, "model_path": "/mnt/nfs_clientshare/JiRackUltra_14b/model.pt", "adapter": "/mnt/nfs_clientshare/JiRackUltra_14b/toolace_lora_adapter.pt", "merged": "/mnt/nfs_clientshare/JiRackUltra_14b/ultra14b_toolace.pt", "max_len": 1024, # [U-D] halve the context to fit "grad_accum": 16, }, "32b": { "module": "JiRackTernaryUltra_32b", "vocab": 152064, "model_path": "/mnt/nfs_clientshare/JiRackUltra_32b/model.pt", "adapter": "/mnt/nfs_clientshare/JiRackUltra_32b/toolace_lora_adapter.pt", "merged": "/mnt/nfs_clientshare/JiRackUltra_32b/ultra32b_toolace.pt", "max_len": 1024, "grad_accum": 32, }, } if SIZE not in SIZES: sys.exit(f"❌ SIZE must be one of {list(SIZES)}, got '{SIZE}'") CFG = SIZES[SIZE] _mod = __import__(CFG["module"], fromlist=["JiRackTransformer", "JiRackConfig"]) JiRackTransformer = _mod.JiRackTransformer JiRackConfig = _mod.JiRackConfig # ========================= EDIT THESE ========================= MODEL_PATH = CFG["model_path"] TOKENIZER_DIR = "CMSManhattan/JiRackPrecisionTokenizer" # [U-C] HF repo or local dir DATASET_PATH = "/mnt/nfs_clientshare/datasets/toolace_sft_jirack_precision.jsonl" ADAPTER_OUT = CFG["adapter"] MERGED_OUT = CFG["merged"] # LoRA LORA_R = 16 LORA_ALPHA = 32 LORA_DROPOUT = 0.05 # Training EPOCHS = 2 LR = 2e-4 # LoRA params EMBED_LR = 5e-5 # new-token embedding rows (gentler) BATCH_SIZE = 1 GRAD_ACCUM = CFG["grad_accum"] MAX_LEN = CFG["max_len"] WARMUP_STEPS = 50 SEED = 42 LAMBDA = 0.0 # 0.0 = full-precision training (recommended: # you're teaching tool-call FORMAT, not # doing QAT -- run the ternarization QAT # scripts separately, AFTER this merge) SAVE_EVERY = 500 # optimizer steps between adapter checkpoints MERGE_AT_END = True OPTIMIZER = "adafactor" # "adamw" or "adafactor" # adafactor: ~2 bytes/param optimizer state # vs AdamW's ~8 -- matters at 14B/32B. # [U-B] FREEZE_8BIT removed on purpose -- see header. # ================================================================ # ------------------------------ LoRA machinery ------------------------------ class LoRALinear(nn.Module): """Wraps a frozen nn.Linear (or BitLinear); adds trainable low-rank A/B.""" def __init__(self, base: nn.Linear, r: int, alpha: int, dropout: float): super().__init__() self.base = base self.r = r self.scale = alpha / r self.lora_A = nn.Parameter(torch.zeros(r, base.in_features)) self.lora_B = nn.Parameter(torch.zeros(base.out_features, r)) nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5)) # B starts at zero -> identity behavior at step 0 self.dropout = nn.Dropout(dropout) if dropout > 0 else nn.Identity() def forward(self, x): out = self.base(x) # [U-A] BitLinear.forward lx = self.dropout(x).to(self.lora_A.dtype) out = out + (lx @ self.lora_A.T @ self.lora_B.T) * self.scale return out @torch.no_grad() def merge_into_base(self): """Fold the LoRA delta into the base weight, in place. The base stays the SAME module object (BitLinear stays BitLinear), so lambda_ buffers and state_dict keys survive untouched.""" delta = (self.lora_B.float() @ self.lora_A.float()) * self.scale self.base.weight.data += delta.to(self.base.weight.dtype) def inject_lora(model, vocab_size): """Replace every nn.Linear (except the vocab-sized head) with LoRALinear. BitLinear subclasses nn.Linear, so the whole ternary backbone gets wrapped -- intended. Embeddings are nn.Embedding, not nn.Linear, so they're skipped here and handled separately by the row-mask logic.""" wrapped = [] for parent_name, parent in list(model.named_modules()): for child_name, child in list(parent.named_children()): if isinstance(child, LoRALinear): continue if isinstance(child, nn.Linear) and child.out_features != vocab_size: setattr(parent, child_name, LoRALinear(child, LORA_R, LORA_ALPHA, LORA_DROPOUT)) full = f"{parent_name}.{child_name}" if parent_name else child_name wrapped.append(full) return wrapped def merge_and_unwrap(model): """Fold LoRA into base weights and restore the original modules, so state_dict() keys match the original checkpoint exactly.""" for parent_name, parent in list(model.named_modules()): for child_name, child in list(parent.named_children()): if isinstance(child, LoRALinear): child.merge_into_base() setattr(parent, child_name, child.base) # ------------------------------ Dataset ------------------------------ def load_dataset(path): convs = [] with open(path) as f: for line in f: line = line.strip() if not line: continue obj = json.loads(line) msgs = obj.get("messages", obj) if isinstance(msgs, list) and any(m.get("role") == "assistant" for m in msgs): convs.append(msgs) return convs def build_example(tokenizer, messages, max_len): """Tokenize a conversation with assistant-only labels. Incremental templating: token span of message i = template(msgs[:i+1]) minus template(msgs[:i]). Labels = ids inside assistant spans, else -100.""" ids, labels = [], [] prev = [] prev_len = 0 for m in messages: prev.append(m) cur = tokenizer.apply_chat_template(prev, tokenize=True, add_generation_prompt=False) span = cur[prev_len:] if m["role"] == "assistant": labels.extend(span) else: labels.extend([-100] * len(span)) ids = cur prev_len = len(cur) if len(ids) >= max_len: break ids = ids[:max_len] labels = labels[:max_len] if all(l == -100 for l in labels): return None return torch.tensor(ids), torch.tensor(labels) # ------------------------------ Training ------------------------------ def main(): random.seed(SEED) torch.manual_seed(SEED) device = "cuda" if torch.cuda.is_available() else "cpu" print(f"🚀 JiRack Ultra {SIZE.upper()} ToolACE LoRA | Device: {device.upper()}") print(f"⚙️ module={CFG['module']} optimizer={OPTIMIZER} " f"MAX_LEN={MAX_LEN} GRAD_ACCUM={GRAD_ACCUM}") tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_DIR) # --- model --- config = JiRackConfig() # [U-C] the config's vocab must match what this script expects for this size assert config.vocab_size == CFG["vocab"], ( f"{CFG['module']}.JiRackConfig has vocab_size={config.vocab_size} but " f"SIZE='{SIZE}' expects {CFG['vocab']} -- wrong module for this size?" ) # [U-C] tokenizer must FIT the padded matrix; never resize assert len(tokenizer) <= config.vocab_size, ( f"tokenizer ({len(tokenizer)}) > padded matrix ({config.vocab_size}) — " f"do NOT resize_token_embeddings, fix the tokenizer instead" ) print(f"✅ Tokenizer fits: {len(tokenizer)} <= {config.vocab_size}") model = JiRackTransformer(config, use_checkpoint=True) # activation ckpt on print(f"📥 Loading {MODEL_PATH} ...") ckpt = torch.load(MODEL_PATH, map_location="cpu", weights_only=False) sd = ckpt["model"] if isinstance(ckpt, dict) and "model" in ckpt else ckpt missing, unexpected = model.load_state_dict(sd, strict=False) real_missing = [k for k in missing if not k.endswith("lambda_")] if real_missing: print(f"⚠️ Missing keys: {real_missing[:10]}") if unexpected: print(f"⚠️ Unexpected keys: {list(unexpected)[:10]}") model = model.to(dtype=torch.bfloat16, device=device) model.set_lambda(LAMBDA) # find embedding module + vocab size embed = None for mod in model.modules(): if isinstance(mod, nn.Embedding): embed = mod break if embed is None: sys.exit("❌ No nn.Embedding found in model") vocab_rows = embed.weight.shape[0] print(f" embedding rows: {vocab_rows}") # --- find the LM head BEFORE wrapping (after injection it'd be hidden) --- head = None for mod in model.modules(): if isinstance(mod, nn.Linear) and mod.out_features == vocab_rows: head = mod break # --- freeze all, inject LoRA --- for p in model.parameters(): p.requires_grad = False wrapped = inject_lora(model, vocab_rows) model = model.to(device) print(f"🧩 LoRA injected into {len(wrapped)} Linear layers " f"(r={LORA_R}, alpha={LORA_ALPHA})") lora_params = [p for n, p in model.named_parameters() if "lora_" in n] for p in lora_params: p.requires_grad = True # --- unfreeze ONLY the JiRack special-token embedding rows --- special_ids = sorted(set(tokenizer.additional_special_tokens_ids or [])) special_ids = [i for i in special_ids if i < vocab_rows] if not special_ids: print("⚠️ No additional_special_tokens found in the tokenizer -- " "training LoRA only, no embedding rows. If you expected the " "JiRack tool-call/robotics tags here, check the tokenizer repo.") embed.weight.requires_grad = True row_mask = torch.zeros(vocab_rows, 1, device=device) for i in special_ids: row_mask[i] = 1.0 embed.weight.register_hook(lambda g: g * row_mask.to(g.dtype)) if special_ids: print(f"🎯 Training embedding rows for {len(special_ids)} special tokens " f"(ids {special_ids[0]}..{special_ids[-1]}), base vocab frozen " f"via grad mask.") # untied lm_head: train the same rows there too (the model can't EMIT a # token whose output row is noise, even with good input embeddings) if head is not None and head.weight is not embed.weight: head.weight.requires_grad = True head.weight.register_hook(lambda g: g * row_mask.to(g.dtype)) print("🎯 LM head is untied -- training the same rows there as well.") elif head is None: print("⚠️ No vocab-sized Linear found -- lm_head not trained.") n_train = sum(p.numel() for p in model.parameters() if p.requires_grad) print(f" trainable params (incl. masked embeds): {n_train/1e6:.1f}M") # --- data --- convs = load_dataset(DATASET_PATH) print(f"📚 {len(convs)} conversations loaded from {DATASET_PATH}") random.shuffle(convs) # --- optimizer --- groups = [{"params": lora_params, "lr": LR}] embed_params = [embed.weight] if head is not None and head.weight is not embed.weight: embed_params.append(head.weight) groups.append({"params": embed_params, "lr": EMBED_LR}) if OPTIMIZER == "adafactor": # relative_step=False + explicit per-group lr so our own cosine # schedule (LambdaLR below) still controls the learning rate. optim = Adafactor(groups, scale_parameter=False, relative_step=False, warmup_init=False, weight_decay=0.0) print("⚙️ Optimizer: Adafactor (relative_step=False, no momentum buffer)") elif OPTIMIZER == "adamw": optim = torch.optim.AdamW(groups, weight_decay=0.0) print("⚙️ Optimizer: AdamW") else: sys.exit(f"❌ Unknown OPTIMIZER '{OPTIMIZER}' -- use 'adamw' or 'adafactor'") total_steps = max(1, (len(convs) * EPOCHS) // (BATCH_SIZE * GRAD_ACCUM)) def lr_lambda(step): if step < WARMUP_STEPS: return step / max(1, WARMUP_STEPS) prog = (step - WARMUP_STEPS) / max(1, total_steps - WARMUP_STEPS) return 0.5 * (1.0 + math.cos(math.pi * min(1.0, prog))) sched = torch.optim.lr_scheduler.LambdaLR(optim, lr_lambda) loss_fn = nn.CrossEntropyLoss(ignore_index=-100) def save_adapter(path): state = {n: p.detach().cpu() for n, p in model.named_parameters() if "lora_" in n} state["__special_ids__"] = torch.tensor(special_ids) if special_ids: state["__embed_rows__"] = embed.weight.detach()[special_ids].cpu() if head is not None and head.weight is not embed.weight: state["__head_rows__"] = head.weight.detach()[special_ids].cpu() torch.save({"size": SIZE, "lora_r": LORA_R, "lora_alpha": LORA_ALPHA, "state": state}, path) print(f"💾 Adapter saved: {path}") # --- loop --- model.train() step, micro, running = 0, 0, 0.0 t0 = time.time() for epoch in range(EPOCHS): for conv in convs: ex = build_example(tokenizer, conv, MAX_LEN) if ex is None: continue ids, labels = ex ids = ids.unsqueeze(0).to(device) labels = labels.unsqueeze(0).to(device) with torch.autocast(device_type=("cuda" if device == "cuda" else "cpu"), dtype=torch.bfloat16): logits = model(ids) loss = loss_fn( logits[:, :-1, :].reshape(-1, logits.size(-1)).float(), labels[:, 1:].reshape(-1)) if torch.isnan(loss) or torch.isinf(loss): print(f"⚠️ NaN/Inf loss at micro-step {micro} — example skipped") optim.zero_grad(set_to_none=True) micro += 1 continue (loss / GRAD_ACCUM).backward() running += loss.item() micro += 1 if micro % GRAD_ACCUM == 0: torch.nn.utils.clip_grad_norm_( [p for p in model.parameters() if p.requires_grad], 1.0) optim.step() sched.step() optim.zero_grad(set_to_none=True) step += 1 if step % 10 == 0: avg = running / (10 * GRAD_ACCUM) running = 0.0 el = time.time() - t0 print(f"epoch {epoch+1} step {step}/{total_steps} " f"loss {avg:.4f} lr {sched.get_last_lr()[0]:.2e} " f"[{el/60:.1f} min]") if step % SAVE_EVERY == 0: save_adapter(ADAPTER_OUT) save_adapter(ADAPTER_OUT) # --- merge --- if MERGE_AT_END: print("🔀 Merging LoRA into base weights ...") model.eval() merge_and_unwrap(model) merged_sd = {k: v.detach().cpu() for k, v in model.state_dict().items()} # drop lambda_ buffers if the original checkpoint didn't carry them orig_keys = set(sd.keys()) merged_sd = {k: v for k, v in merged_sd.items() if k in orig_keys or not k.endswith("lambda_")} extra = set(merged_sd.keys()) - orig_keys missing2 = orig_keys - set(merged_sd.keys()) if extra: print(f"⚠️ Keys not in original ckpt (kept): {list(extra)[:8]}") if missing2: print(f"⚠️ Original keys absent in merged (check!): {list(missing2)[:8]}") torch.save(merged_sd, MERGED_OUT) print(f"✅ Merged checkpoint saved: {MERGED_OUT}") print(f" Next: point your chat script at it, verify tool tags are " f"emitted, THEN run the ternarization QAT script for {SIZE}.") if __name__ == "__main__": main()