JiRackUltra_1b / qat_ultra_1b_for_TQ_2.py
kgrabko's picture
Rename qat_exmple_for_QT_2.py to qat_ultra_1b_for_TQ_2.py
4ea572a verified
Raw History Blame Contribute Delete
16.6 kB
#%%writefile train_ultra_1b_bf16_ada_warmup.py
# =============================================================================
# COPYRIGHT © 2025-2026 Konstantin Vladimirovich Grabko. ALL RIGHTS RESERVED.
# CMS Manhattan JiRack Technology — PATENT PENDING
# train_ultra_1b_bf16_ada_warmup.py
# =============================================================================
# QAT training with lambda warmup — JiRack Ultra 1B edition.
# Adapted from train_openassist_blackwell_96gb_bf16_ada_warmup_v2.py (10B).
#
# Changes vs the 10B script:
# [U-1] import from JiRackTernaryUltra_1b (user's class: JiRackTransformer,
# JiRackConfig — same set_lambda/get_lambda API as the 10B fixed file)
# [U-2] base weights come from the OFFICIAL HF checkpoint
# deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B via model.load_hf_state_dict()
# — no migrate_checkpoint.py step exists/needed for 1B.
# Optional: BASE_CHECKPOINT .pt still supported if you have one.
# [U-3] gradient checkpointing: passed via constructor (use_checkpoint=True);
# the per-block flag is also set for parity, blocks read it at forward.
# [U-4] 1.5B fits GPUs the 10B never could — batch/accum defaults raised;
# tune for your card (T4 16GB: BATCH_SIZE=4–8 at seq 1024).
# [U-5] vocab_size read from config attr (config.vocab_size, same name).
# [U-6] tokenizer sanity gate: CMSManhattan/JiRackPrecisionTokenizer must
# fit the padded matrix (151,779 <= 151,936) — assert, NEVER resize.
# Gate is optional (SKIP_TOKENIZER_CHECK) since prepared shards are
# already tokenized; it just catches wrong-family data early.
# All 10B fixes preserved: [T-3] resume restores global_step (legacy
# checkpoints reconstruct it by inverting the sigmoid), [T-4] Adafactor
# state saved/restored, [T-5] lambda updated per accumulation window,
# [T-7] val loss logged with its lambda, atomic mid-shard autosave.
# =============================================================================
import os
# Must be set BEFORE torch initializes CUDA.
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import torch
import glob
import re
import gc
import math
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
from torch.nn.utils.rnn import pad_sequence
from tqdm import tqdm
from transformers.optimization import Adafactor
# [U-1] user's model class
from JiRackTernaryUltra_1b import JiRackTransformer, JiRackConfig
# ========================= НАСТРОЙКИ =========================
DATA_DIR = "/content/prepared_sft_data"
OUTPUT_DIR = "/content/JiRackUltra_1B_Checkpoints"
# [U-2] Base weights: JiRack Ultra 1B published checkpoint (Qwen2-compatible
# config: vocab 151936, hidden 1536, 28L, 12/2 heads — matches JiRackConfig,
# loads via the same load_hf_state_dict mapping as the original distill).
HF_BASE_MODEL = "CMSManhattan/JiRackUltra_1b"
# Optional local .pt to use INSTEAD of HF (set to a path or leave None).
BASE_CHECKPOINT = None # e.g. "/content/ultra1b_base.pt"
# [U-4] 1.5B is small — these are safe on a 16GB card at seq<=1024;
# raise BATCH_SIZE / lower GRAD_ACCUM to taste on bigger GPUs.
BATCH_SIZE = 4
GRAD_ACCUM = 4
LR = 4e-5
VAL_RATIO = 0.05
# === Lambda Warmup Settings ===
# Same gentle sigmoid as 10B; measured in MICRO-batches.
# 5000 micro-batches = 1250 optimizer steps at GRAD_ACCUM=4.
LAMBDA_WARMUP_STEPS = 5000
MAX_LAMBDA = 1.0
SAVE_OPTIMIZER = True # [T-4]
AUTOSAVE_EVERY = 1000 # mid-shard autosave, 0 = off
# [U-6] tokenizer — CMSManhattan/JiRackPrecisionTokenizer, kept loaded for
# real use (pad_token_id below), not just a throwaway fit-check.
TOKENIZER_REPO = "CMSManhattan/JiRackPrecisionTokenizer"
# =============================================================
os.makedirs(OUTPUT_DIR, exist_ok=True)
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
print("🚀 Loading JiRack Ultra 1B Ternary + Lambda Warmup...")
config = JiRackConfig()
# [U-6] load once, keep it — assert fit against the padded matrix,
# NEVER resize_token_embeddings (151,779 < 151,936: shrinking would
# corrupt the matrix, despite the "Must resize" note on the tokenizer's
# own HF card — that note only applies when tokenizer vocab > model vocab).
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_REPO)
assert len(tokenizer) <= config.vocab_size, (
f"tokenizer ({len(tokenizer)}) > padded matrix ({config.vocab_size}) — "
f"do NOT resize_token_embeddings, fix the tokenizer/data instead"
)
print(f"✅ Tokenizer fits padded matrix: {len(tokenizer)} <= {config.vocab_size}")
# The tokenizer card shows inconsistent pad/eos ids in different sections
# (pad/eos 151643 in one block, eos 151645 in the "Tesr Tokenizer size" test
# output). Trust the LOADED object, not the card text, and fail loudly if
# pad_token_id is unset rather than silently defaulting to 0.
PAD_ID = tokenizer.pad_token_id
if PAD_ID is None:
PAD_ID = tokenizer.eos_token_id
print(f"⚠️ tokenizer.pad_token_id is None, falling back to eos_token_id={PAD_ID}")
assert PAD_ID is not None, "tokenizer has neither pad_token_id nor eos_token_id set"
print(f"✅ Using pad_token_id={PAD_ID} (eos_token_id={tokenizer.eos_token_id})")
# [U-3] use_checkpoint via constructor (blocks capture it at build time)
model = JiRackTransformer(config, use_checkpoint=True)
model.to("cuda")
for block in model.blocks:
block.use_checkpoint = True # parity with 10B script; harmless duplicate
print("✅ Gradient checkpointing enabled")
optimizer = Adafactor(
model.parameters(),
lr=LR,
eps=(1e-30, 1e-3),
clip_threshold=1.0,
decay_rate=-0.8,
weight_decay=0.0001,
scale_parameter=False,
relative_step=False,
warmup_init=False,
)
criterion = nn.CrossEntropyLoss(ignore_index=-100)
# ==================== LAMBDA SCHEDULE ====================
def lambda_schedule(step: int) -> float:
"""Sigmoid ramp over LAMBDA_WARMUP_STEPS micro-batches."""
if step >= LAMBDA_WARMUP_STEPS:
return MAX_LAMBDA
return MAX_LAMBDA / (
1.0 + math.exp(-10 * (step - LAMBDA_WARMUP_STEPS / 2) / LAMBDA_WARMUP_STEPS)
)
def invert_lambda_schedule(lam: float) -> int:
"""Reconstruct global_step from a saved lambda (legacy checkpoints).
Inverse of the sigmoid above."""
if lam >= MAX_LAMBDA * 0.999:
return LAMBDA_WARMUP_STEPS
if lam <= 1e-6:
return 0
p = lam / MAX_LAMBDA
x = -math.log(1.0 / p - 1.0) # logit
return int(round(x * LAMBDA_WARMUP_STEPS / 10 + LAMBDA_WARMUP_STEPS / 2))
# ==================== CHECKPOINT LOAD ====================
def load_any_checkpoint(path, model, optimizer):
"""Handles both new-format dicts and legacy plain state_dicts.
Returns restored global_step."""
ckpt = torch.load(path, map_location="cpu", weights_only=True)
if isinstance(ckpt, dict) and "model" in ckpt:
# [T-3] new format
missing, unexpected = model.load_state_dict(ckpt["model"], strict=False)
assert not unexpected, unexpected
assert all(k.endswith("lambda_") for k in missing), missing
if SAVE_OPTIMIZER and "optimizer" in ckpt and ckpt["optimizer"] is not None:
try:
optimizer.load_state_dict(ckpt["optimizer"])
print("✅ Optimizer state restored")
except Exception as e:
print(f"⚠️ Optimizer state not restored ({e}); continuing fresh")
step = int(ckpt.get("global_step", 0))
print(f"✅ Resumed (new format): global_step={step}, "
f"lambda={ckpt.get('lambda', 'n/a')}")
return step
# legacy: plain state_dict
missing, unexpected = model.load_state_dict(ckpt, strict=False)
assert not unexpected, unexpected
assert all(k.endswith("lambda_") for k in missing), missing
lam = model.get_lambda()
step = invert_lambda_schedule(lam)
print(f"✅ Resumed (legacy format): lambda={lam:.4f} -> "
f"reconstructed global_step={step}")
return step
def load_hf_base(model):
"""[U-2] Pull the official 1.5B distill and map it in with the model's
own load_hf_state_dict (strict; tolerates only missing lambda_)."""
from transformers import AutoModelForCausalLM
print(f"⬇️ Loading HF base: {HF_BASE_MODEL}")
hf = AutoModelForCausalLM.from_pretrained(HF_BASE_MODEL, torch_dtype=torch.float32)
real_missing, unexpected = model.load_hf_state_dict(hf.state_dict(), strict=True)
del hf
gc.collect()
torch.cuda.empty_cache()
return 0 # fresh QAT run starts at global_step 0
# Shard-named checkpoints (define which shards are already done)...
checkpoints = sorted(
glob.glob(os.path.join(OUTPUT_DIR, "jirack_ultra1b_data_*.pt")),
key=lambda x: int(re.search(r"data_(\d+)", x).group(1)),
)
# ...but for MODEL STATE, resume from whichever .pt is newest on disk.
all_ckpts = glob.glob(os.path.join(OUTPUT_DIR, "*.pt"))
global_step = 0
if all_ckpts:
LATEST_CKPT = max(all_ckpts, key=os.path.getmtime)
print(f"📦 Resuming from (newest on disk): {LATEST_CKPT}")
global_step = load_any_checkpoint(LATEST_CKPT, model, optimizer)
elif BASE_CHECKPOINT is not None:
print(f"📦 Loading base checkpoint: {BASE_CHECKPOINT}")
assert os.path.exists(BASE_CHECKPOINT), f"{BASE_CHECKPOINT} not found"
global_step = load_any_checkpoint(BASE_CHECKPOINT, model, optimizer)
else:
# [U-2] default path for 1B: official HF weights, QAT from step 0
global_step = load_hf_base(model)
# Model must be on GPU in bf16-friendly state after any load path.
model.to("cuda")
# [T-3] lambda follows global_step from here on.
model.set_lambda(lambda_schedule(global_step))
print(f"🔧 Warmup: {LAMBDA_WARMUP_STEPS} steps | starting at "
f"step={global_step}, lambda={model.get_lambda():.4f}")
# ========================= DATASET =========================
class ShardDataset(Dataset):
def __init__(self, data_list):
self.data = data_list
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
def collate_fn(batch):
# [U-6] pad with the tokenizer's real pad_token_id, not a hardcoded 0 —
# 0 may already be a meaningful token in this vocab.
input_ids = pad_sequence(
[item["input_ids"] for item in batch], batch_first=True, padding_value=PAD_ID
)
attention_mask = pad_sequence(
[item.get("attention_mask", torch.ones_like(item["input_ids"]))
for item in batch],
batch_first=True, padding_value=0,
)
labels = input_ids.clone()
labels[attention_mask == 0] = -100
return {"input_ids": input_ids, "labels": labels}
# ========================= CHECKPOINT SAVE =========================
def save_checkpoint(path, model, optimizer, global_step, lam):
# bf16 on disk for economy (fp32 master precision lost across restarts only)
model_sd = {k: v.detach().to(torch.bfloat16).cpu()
for k, v in model.state_dict().items()}
ckpt = {
"model": model_sd,
"optimizer": optimizer.state_dict() if SAVE_OPTIMIZER else None,
"global_step": global_step,
"lambda": lam,
}
torch.save(ckpt, path)
# ========================= TRAINING =========================
all_shards = sorted(
glob.glob(f"{DATA_DIR}/sft_data_*.pt"),
key=lambda x: int(re.search(r"sft_data_(\d+)", x).group(1)),
)
last_done = -1 # -1 = no shards processed yet (shard numbering starts at 0!)
if checkpoints:
m = re.search(r"data_(\d+)", checkpoints[-1])
last_done = int(m.group(1)) if m else -1
for shard_path in all_shards:
shard_name = os.path.basename(shard_path)
shard_num = int(re.search(r"sft_data_(\d+)", shard_name).group(1))
if shard_num <= last_done:
print(f"⏭ Skipping already processed: {shard_name}")
continue
print(f"\n🔥 Starting shard: {shard_name}")
raw_shard_data = torch.load(shard_path, map_location="cpu", weights_only=False)
val_size = int(len(raw_shard_data) * VAL_RATIO)
train_size = len(raw_shard_data) - val_size
train_data, val_data = torch.utils.data.random_split(
raw_shard_data, [train_size, val_size],
generator=torch.Generator().manual_seed(42),
)
train_loader = DataLoader(
ShardDataset(train_data), batch_size=BATCH_SIZE, shuffle=True,
collate_fn=collate_fn, pin_memory=True,
)
model.train()
pbar = tqdm(train_loader, desc=f"Shard {shard_num}", dynamic_ncols=True)
optimizer.zero_grad()
lambda_value = lambda_schedule(global_step)
model.set_lambda(lambda_value)
for step, batch in enumerate(pbar):
# [T-5] update lambda only at accumulation-window boundaries.
if step % GRAD_ACCUM == 0:
lambda_value = lambda_schedule(global_step)
model.set_lambda(lambda_value)
input_ids = batch["input_ids"].to("cuda", non_blocking=True)
labels = batch["labels"].to("cuda", non_blocking=True)
with torch.amp.autocast("cuda", dtype=torch.bfloat16):
logits = model(input_ids)
if isinstance(logits, tuple):
logits = logits[0]
loss = criterion(
logits[..., :-1, :].reshape(-1, config.vocab_size),
labels[..., 1:].reshape(-1),
)
loss = loss / GRAD_ACCUM
if torch.isnan(loss) or torch.isinf(loss):
print(f"\n⚠️ NaN/Inf loss at step {global_step} "
f"(lambda={lambda_value:.4f}) — window dropped")
optimizer.zero_grad(set_to_none=True)
torch.cuda.empty_cache()
global_step += 1
continue
loss.backward()
if (step + 1) % GRAD_ACCUM == 0:
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
optimizer.zero_grad()
if step % 10 == 0:
pbar.set_postfix({
"loss": f"{loss.item() * GRAD_ACCUM:.4f}",
"lambda": f"{lambda_value:.4f}",
"gstep": global_step,
})
# Mid-shard autosave (atomic: tmp -> rename), at accumulation
# boundaries only (grads flushed, state clean).
if (AUTOSAVE_EVERY > 0 and global_step > 0
and global_step % AUTOSAVE_EVERY == 0
and (step + 1) % GRAD_ACCUM == 0):
autosave_path = os.path.join(OUTPUT_DIR, "autosave_latest.pt")
tmp_path = autosave_path + ".tmp"
save_checkpoint(tmp_path, model, optimizer, global_step, lambda_value)
os.replace(tmp_path, autosave_path)
pbar.write(f"💾 autosave @ gstep={global_step}, "
f"lambda={lambda_value:.4f}")
global_step += 1
# ==================== Validation ====================
print("🧪 Validating...")
model.eval()
total_val_loss = 0.0
val_steps = 0
val_loader = DataLoader(
ShardDataset(val_data), batch_size=BATCH_SIZE, shuffle=False,
collate_fn=collate_fn, pin_memory=True,
)
with torch.no_grad():
for batch in tqdm(val_loader, desc="Validating", leave=False):
input_ids = batch["input_ids"].to("cuda", non_blocking=True)
labels = batch["labels"].to("cuda", non_blocking=True)
with torch.amp.autocast("cuda", dtype=torch.bfloat16):
logits = model(input_ids)
if isinstance(logits, tuple):
logits = logits[0]
v_loss = criterion(
logits[..., :-1, :].reshape(-1, config.vocab_size),
labels[..., 1:].reshape(-1),
)
if not (torch.isnan(v_loss) or torch.isinf(v_loss)):
total_val_loss += v_loss.item()
val_steps += 1
avg_val_loss = total_val_loss / val_steps if val_steps > 0 else float("inf")
# [T-7] val loss only comparable at the SAME lambda
print(f"📊 Shard {shard_num} — Val Loss: {avg_val_loss:.4f} "
f"@ lambda={lambda_value:.4f} (gstep={global_step})")
# ==================== Save ====================
save_path = os.path.join(OUTPUT_DIR, f"jirack_ultra1b_data_{shard_num}.pt")
save_checkpoint(save_path, model, optimizer, global_step, lambda_value)
print(f"💾 Saved: {save_path}")
del raw_shard_data
torch.cuda.empty_cache()
gc.collect()
print("🏁 Training finished with Lambda Warmup!")