FST_code / src /lmr /training /trainer_bert.py
jasonfan's picture
2026-03-19
3b2d368 verified
Raw
History Blame
26.4 kB
import math
import traceback
import time
from pathlib import Path
from contextlib import nullcontext
from tqdm import tqdm
import torch
import torch.nn.functional as F
from torch.amp import GradScaler, autocast
from torch.utils.data import DataLoader
from transformers import get_cosine_schedule_with_warmup
from lmr.ddp import setup_ddp, cleanup_ddp, initialize_model_ddp, unwrap_model, initialize_samplers_ddp
from lmr.utils.logger import Logger
from safetensors.torch import load_file
from pathlib import Path
import json
import os
from huggingface_hub import HfApi, upload_folder
def _get_hf_tokenizer(tok):
"""Unwrap common tokenizer wrappers to a HF tokenizer."""
if tok is None:
return None
if hasattr(tok, "save_pretrained"):
return tok
for attr in ["hf", "hf_tokenizer", "tokenizer", "base_tokenizer", "_tokenizer"]:
if hasattr(tok, attr):
inner = getattr(tok, attr)
if hasattr(inner, "save_pretrained"):
return inner
return None
class Bert_Trainer:
def __init__(self, training_config, model, tokenizer, splits, checkpointing, samplers=None, device=None):
self.training_config = training_config
self.model = model
self.tokenizer = tokenizer
self.splits = splits
self.checkpointing = checkpointing
self.samplers = samplers
self.device = device
# Special tokens
self.pad_token_id = getattr(self.tokenizer, "pad_token_id", None)
if self.pad_token_id is None:
self.pad_token_id = 0
self.mask_token_id = getattr(self.tokenizer, "mask_token_id", None)
if self.mask_token_id is None:
raise ValueError("Tokenizer must have a mask_token_id for MLM (e.g., '[MASK]').")
self.eos_token_id = getattr(self.tokenizer, "eos_token_id", None)
# precision
self.autocast_dtype = getattr(torch, self.training_config.precision)
# ddp
self.use_ddp = False
self.rank = 0
self.world_size = 1
self.debug_prompts = [
"Question: What is 15 + 32?\nAnswer:",
"Question: There are 5 birds on a tree. 2 fly away. How many are left?\nSolution:",
]
def _is_lora_param(self, param_name):
lora_indicators = ['lora_A', 'lora_B', 'lora_dropout']
return any(indicator in param_name for indicator in lora_indicators)
def save_as_hf(
self,
save_dir: str,
repo_id: str | None = None,
private: bool = True,
push_to_hub: bool = True,
push_config_fixes: bool = True,
commit_message: str = "Upload model from trainer.save_as_hf",
):
"""
1) Save model/tokenizer to HF format
2) (Optional) Automatically upload to Hugging Face Hub
repo_id example: "jf381/fst_353M_bert_medium"
"""
# -----------------------------
# DDP safety
# -----------------------------
if self.use_ddp and self.rank != 0:
return
save_path = Path(save_dir)
save_path.mkdir(parents=True, exist_ok=True)
model_to_save = unwrap_model(self.model)
model_to_save.eval()
# -----------------------------
# 1) Save model
# -----------------------------
model_to_save.save_pretrained(
str(save_path),
safe_serialization=True
)
# -----------------------------
# 2) Save tokenizer (unwrap)
# -----------------------------
hf_tok = _get_hf_tokenizer(self.tokenizer)
if hf_tok is not None:
hf_tok.save_pretrained(str(save_path))
else:
# fallback metadata (won't be fully loadable)
(save_path / "tokenizer_wrapper.json").write_text(
json.dumps(
{
"warning": "Tokenizer has no save_pretrained(); "
"HF tokenizer files missing.",
"tokenizer_class": self.tokenizer.__class__.__name__,
},
indent=2,
)
)
# -----------------------------
# 3) Patch config.json (optional)
# -----------------------------
if push_config_fixes:
cfg_file = save_path / "config.json"
if cfg_file.exists():
cfg = json.loads(cfg_file.read_text())
cfg.setdefault("architectures", [model_to_save.__class__.__name__])
cfg.setdefault("model_type", cfg.get("model_type", "bert"))
cfg_file.write_text(json.dumps(cfg, indent=2))
model_to_save.train()
Logger.log(f"📦 HF model saved locally at: {save_path}")
# -----------------------------
# 4) Push to Hugging Face Hub
# -----------------------------
if push_to_hub:
if repo_id is None:
raise ValueError("push_to_hub=True but repo_id is None")
token = os.getenv("HF_TOKEN", None) # optional if huggingface-cli login used
api = HfApi()
api.create_repo(
repo_id=repo_id,
repo_type="model",
private=private,
exist_ok=True,
token=token,
)
upload_folder(
folder_path=str(save_path),
repo_id=repo_id,
repo_type="model",
token=token,
commit_message=commit_message,
)
Logger.log(f"🚀 Uploaded to Hugging Face Hub: https://huggingface.co/{repo_id}")
# =========================================================================
# Debug helpers
# =========================================================================
def _log_batch_samples(self, batch, title="SAMPLE CHECK"):
if self.rank != 0:
return
num_show = min(batch.size(0), 2)
print(f"\n{'='*20} {title} (First {num_show} samples) {'='*20}")
for i in range(num_show):
seq_ids = batch[i].tolist()
display_ids = [x for x in seq_ids if x != self.pad_token_id and x != -100]
try:
text = self.tokenizer.decode(display_ids)
except Exception as e:
text = f"[Decode Error: {e}]"
print(f"[Sample {i}]")
print(f" Tokens: {display_ids[:40]} ...")
print(f" Text: {text[:200]} ...")
print("-" * 40)
print(f"{'='*60}\n")
def load_only_model_weights(self, checkpoint_path, map_location="cpu", strict=True, verbose=True):
path_obj = Path(checkpoint_path)
if not path_obj.exists():
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
model_state = {}
is_sharded = False
# --- 1. Load the Checkpoint (Sharded or Single) ---
if path_obj.is_dir():
index_file = path_obj / "model.safetensors.index.json"
if index_file.exists():
if verbose: print(f"🔹 Detected sharded safetensors folder: {path_obj}")
is_sharded = True
import json
with open(index_file, 'r') as f:
index_data = json.load(f)
weight_map = index_data.get("weight_map", {})
shards = set(weight_map.values())
for shard_name in shards:
shard_path = path_obj / shard_name
shard_weights = load_file(str(shard_path), device=str(map_location))
model_state.update(shard_weights)
else:
# Fallback to single file in dir
possible = list(path_obj.glob("*.safetensors")) + list(path_obj.glob("*.pt"))
if not possible: raise FileNotFoundError(f"No weights in {path_obj}")
path_obj = possible[0]
if not is_sharded and path_obj.is_file():
if path_obj.suffix == ".safetensors":
model_state = load_file(str(path_obj), device=str(map_location))
else:
if verbose: print(f"🔹 Loading pickle (.pt): {path_obj}")
ckpt = torch.load(str(path_obj), map_location=map_location)
model_state = ckpt.get("model", ckpt.get("state_dict", ckpt))
# --- 2. Normalize Keys for Matching ---
ckpt_keys_map = {} # cleaned_key -> original_ckpt_key
for k in model_state.keys():
clean_k = k.replace("module.", "").replace("_orig_mod.", "").replace("model.", "")
ckpt_keys_map[clean_k] = k
load_target = unwrap_model(self.model)
target_state = load_target.state_dict()
filtered_state = {}
missing_in_ckpt = []
size_mismatches = []
for k_target, v_target in target_state.items():
k_target_clean = k_target.replace("module.", "").replace("_orig_mod.", "").replace("model.", "")
if k_target_clean in ckpt_keys_map:
real_ckpt_key = ckpt_keys_map[k_target_clean]
v_ckpt = model_state[real_ckpt_key]
if v_ckpt.shape == v_target.shape:
filtered_state[k_target] = v_ckpt
else:
size_mismatches.append(f"{k_target} (ckpt: {v_ckpt.shape}, target: {v_target.shape})")
else:
missing_in_ckpt.append(k_target)
try:
msg = load_target.load_state_dict(filtered_state, strict=False)
if verbose:
print(f"✅ Weights loaded.")
print(f" - Matched keys: {len(filtered_state)}")
print(f" - Missing keys: {len(missing_in_ckpt)}")
if len(missing_in_ckpt) > 0:
real_missing = [k for k in missing_in_ckpt if not self._is_lora_param(k)]
if real_missing:
print(f"⚠️ Real Missing (non-LoRA): {len(real_missing)} (e.g. {real_missing[:3]})")
print(f" (Target clean key example: {k_target_clean})")
print(f" (Ckpt clean key example: {list(ckpt_keys_map.keys())[0]})")
return msg
except Exception as e:
raise RuntimeError(f"Failed to load model weights: {e}")
# =========================================================================
# MLM masking
# =========================================================================
def _mask_tokens(self, inputs: torch.Tensor):
"""
Standard BERT MLM:
- 15% selected for prediction
- 80% -> [MASK], 10% -> random, 10% -> keep
Returns: inputs_masked, labels (-100 for non-masked)
"""
device = inputs.device
labels = inputs.clone()
# special tokens mask
try:
special_tokens_mask = [
self.tokenizer.get_special_tokens_mask(x, already_has_special_tokens=True)
for x in inputs.tolist()
]
special_tokens_mask = torch.tensor(special_tokens_mask, dtype=torch.bool, device=device)
except Exception:
special_tokens_mask = inputs.eq(self.pad_token_id)
probability_matrix = torch.full(labels.shape, 0.15, device=device)
probability_matrix.masked_fill_(special_tokens_mask, value=0.0)
masked_mask = torch.bernoulli(probability_matrix).bool()
labels[~masked_mask] = -100
inputs_masked = inputs.clone()
rand_for_each = torch.rand(labels.shape, device=device)
mask_token_mask = masked_mask & (rand_for_each < 0.8)
random_token_mask = masked_mask & (rand_for_each >= 0.8) & (rand_for_each < 0.9)
if mask_token_mask.any():
inputs_masked[mask_token_mask] = self.mask_token_id
if random_token_mask.any():
try:
vocab_size = self.tokenizer.vocab_size
except Exception:
vocab_size = len(self.tokenizer.get_vocab())
rand_tokens = torch.randint(low=0, high=vocab_size, size=labels.shape, device=device)
inputs_masked[random_token_mask] = rand_tokens[random_token_mask]
return inputs_masked, labels
# =========================================================================
# Core loss + accuracy
# =========================================================================
def _forward_logits(self, input_ids: torch.Tensor):
outputs = unwrap_model(self.model)(
input_ids=input_ids,
attention_mask=(input_ids != self.pad_token_id).long()
)
return outputs.logits if hasattr(outputs, "logits") else outputs[0]
def _step_loss_and_acc(self, batch: torch.Tensor):
"""
Returns:
loss (scalar)
correct (masked positions correct count)
total (masked positions total count)
"""
inputs = batch.to(self.device, non_blocking=True)
inputs_masked, labels = self._mask_tokens(inputs)
with autocast(device_type="cuda", dtype=self.autocast_dtype):
logits = self._forward_logits(inputs_masked)
loss = F.cross_entropy(
logits.view(-1, logits.size(-1)),
labels.view(-1),
ignore_index=-100
)
with torch.no_grad():
preds = torch.argmax(logits, dim=-1)
mask = labels.ne(-100)
correct = (preds.eq(labels) & mask).sum()
total = mask.sum()
return loss, correct, total
def _step_loss(self, batch: torch.Tensor):
loss, _, _ = self._step_loss_and_acc(batch)
return loss
# =========================================================================
# DDP helpers
# =========================================================================
def _ddp_barrier(self):
if self.use_ddp:
torch.distributed.barrier()
def _reduce(self, item):
if self.use_ddp:
torch.distributed.all_reduce(item, op=torch.distributed.ReduceOp.SUM)
# =========================================================================
# Setup
# =========================================================================
def _get_dataloader(self, split_name):
num_workers = 1 if split_name == "validation" else max(1, self.training_config.num_workers - 1)
return DataLoader(
self.splits[split_name],
batch_size=self.training_config.batch_size,
num_workers=num_workers,
shuffle=(split_name == "train" and self.samplers is None),
sampler=None if self.samplers is None else self.samplers[split_name],
pin_memory=True,
drop_last=True
)
def _initialize_optimizer(self):
# full params training (your LoRA logic removed for brevity; add back if needed)
self.optimizer = torch.optim.AdamW(
self.model.parameters(),
lr=self.training_config.lr,
betas=self.training_config.betas,
weight_decay=self.training_config.weight_decay
)
self.checkpointing.optimizer = self.optimizer
def _initialize_scheduler(self):
self.scheduler = get_cosine_schedule_with_warmup(
optimizer=self.optimizer,
num_warmup_steps=self.training_config.warmup_steps,
num_training_steps=self.steps_per_epoch * self.training_config.max_epochs
)
self.checkpointing.scheduler = self.scheduler
def _initialize_scaler(self):
if self.autocast_dtype == torch.float16:
self.scaler = GradScaler("cuda")
else:
self.scaler = None
self.checkpointing.scaler = self.scaler
def _setup_training(self):
# dataloaders
self.train_dataloader = self._get_dataloader("train")
self.validation_dataloader = self._get_dataloader("validation")
# grad accum
if self.training_config.use_grad_accum and self.training_config.grad_accum_steps == "auto":
tokens_per_model_step = self.training_config.batch_size * self.model.config.max_seq_len * self.world_size
self.grad_accum_steps = max(1, self.training_config.tokens_per_step // tokens_per_model_step)
elif self.training_config.use_grad_accum:
self.grad_accum_steps = int(self.training_config.grad_accum_steps)
else:
self.grad_accum_steps = 1
self.steps_per_epoch = len(self.train_dataloader) // self.grad_accum_steps
self.tokens_per_batch = self.training_config.batch_size * self.model.config.max_seq_len
self.tokens_per_step = self.grad_accum_steps * self.tokens_per_batch * self.world_size
self.tokens_per_epoch = self.tokens_per_step * self.steps_per_epoch
# try resume states
try:
self.checkpointing.load_model_states("recent")
except Exception:
pass
# self.save_as_hf('/work/jf381/checkpoints/transformer_353M_Bert','jasonfan/transformer_353M_Bert',False)
# device + compile
self.device = torch.device(f"cuda:{self.rank}")
if getattr(self.training_config, "compile", False):
self.model = torch.compile(self.model, mode=self.training_config.compile_mode)
self.model.to(self.device)
# 5. Explicit Resume (CLI override)
resume_path = getattr(self.training_config, "resume_checkpoint_path", None)
if resume_path:
Logger.log(f"🔄 Forcing resume from: {resume_path}")
self.load_only_model_weights(resume_path, map_location="cpu", strict=False)
# ddp wrap
if self.use_ddp:
self.model = initialize_model_ddp(self.model, self.rank)
self.model.train()
# optimizer/scheduler/scaler
self._initialize_optimizer()
self._initialize_scheduler()
self._initialize_scaler()
# load optimizer states
try:
self.checkpointing.load_training_states("recent")
except Exception:
pass
# =========================================================================
# Validation
# =========================================================================
def _validate(self):
self.model.eval()
loss_sum = torch.tensor(0.0, device=self.device)
token_count = torch.tensor(0, device=self.device, dtype=torch.long)
correct_sum = torch.tensor(0, device=self.device, dtype=torch.long)
masked_count = torch.tensor(0, device=self.device, dtype=torch.long)
with torch.no_grad():
for batch in tqdm(self.validation_dataloader, desc="Validating", leave=False):
loss, correct, total = self._step_loss_and_acc(batch)
loss = loss.detach()
bsz, seq_len = batch.size(0), batch.size(1)
tokens = bsz * seq_len
loss_sum += loss * tokens
token_count += tokens
correct_sum += correct
masked_count += total
# ddp reduce
self._reduce(loss_sum)
self._reduce(token_count)
self._reduce(correct_sum)
self._reduce(masked_count)
self.model.train()
val_loss = (loss_sum / token_count).item()
val_acc = (correct_sum.float() / masked_count.clamp_min(1).float()).item()
return val_loss, val_acc
# =========================================================================
# Training loop
# =========================================================================
def _calculate_training_tokens(self, epoch, step):
return epoch * self.tokens_per_epoch + step * self.tokens_per_step
def _train(self):
self._setup_training()
# data sanity check
try:
first_batch = next(iter(self.train_dataloader))
self._log_batch_samples(first_batch, title="TRAINING START DATA CHECK")
except StopIteration:
Logger.log("⚠️ Train dataloader is empty!")
start_epoch = self.checkpointing.epoch
start_step = self.checkpointing.step
tokens_trained = self.checkpointing.tokens_trained
resume = start_step != 0
if self.rank == 0:
Logger.log(f"{'Resuming' if resume else 'Starting'} training | Device: {self.device} | DDP: {self.use_ddp}")
mr_step_loss = self.checkpointing.train_loss
mr_val_loss = self.checkpointing.val_loss
mr_val_acc = getattr(self.checkpointing, "val_acc", None) # may be absent
for epoch in range(start_epoch, self.training_config.max_epochs):
if self.train_dataloader.sampler is not None and hasattr(self.train_dataloader.sampler, "set_epoch"):
self.train_dataloader.sampler.set_epoch(epoch)
pbar = tqdm(total=self.steps_per_epoch, desc=f"Epoch {epoch}") if self.rank == 0 else None
step_loss_accum = 0.0
# (optional) train acc accumulators for logging
train_correct = torch.tensor(0, device=self.device, dtype=torch.long)
train_total = torch.tensor(0, device=self.device, dtype=torch.long)
for micro_step, batch in enumerate(self.train_dataloader):
step = micro_step // self.grad_accum_steps
# if micro_step >= 40000:
# return
is_update_step = ((micro_step + 1) % self.grad_accum_steps == 0)
if step >= self.steps_per_epoch:
break
# resume skip
if resume and step < start_step:
if pbar is not None and is_update_step:
pbar.update(1)
self.scheduler.step()
continue
elif resume:
resume = False
sync_ctx = self.model.no_sync() if (self.use_ddp and not is_update_step) else nullcontext()
with sync_ctx:
loss, correct, total = self._step_loss_and_acc(batch)
loss = loss / self.grad_accum_steps
step_loss_accum += loss.item()
# accumulate train acc stats
train_correct += correct
train_total += total
if self.scaler is not None:
self.scaler.scale(loss).backward()
else:
loss.backward()
if not is_update_step:
continue
# optimizer step
if self.scaler is not None:
self.scaler.unscale_(self.optimizer)
torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=1.0)
self.scaler.step(self.optimizer)
self.scaler.update()
else:
torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=1.0)
self.optimizer.step()
self.optimizer.zero_grad(set_to_none=True)
self.scheduler.step()
tokens_trained = self._calculate_training_tokens(epoch, step + 1)
mr_step_loss = step_loss_accum
step_loss_accum = 0.0
# validation checkpoint
if self.training_config.validation_steps is not None and (step + 1) % self.training_config.validation_steps == 0:
self._ddp_barrier()
mr_val_loss, mr_val_acc = self._validate()
# reduce train acc across ddp before saving/logging (optional)
self._reduce(train_correct)
self._reduce(train_total)
train_acc = (train_correct.float() / train_total.clamp_min(1).float()).item()
# reset train acc counters after checkpoint
train_correct.zero_()
train_total.zero_()
self.checkpointing.save_checkpoint(
epoch=epoch,
step=step + 1,
train_loss=mr_step_loss,
val_loss=mr_val_loss,
tokens_trained=tokens_trained,
val_acc=mr_val_acc,
train_acc=train_acc,
)
self._ddp_barrier()
if pbar is not None:
if mr_val_acc is None:
pbar.set_postfix(loss=f"{mr_step_loss:.4f}", val_loss=f"{mr_val_loss:.4f}")
else:
pbar.set_postfix(loss=f"{mr_step_loss:.4f}", val_loss=f"{mr_val_loss:.4f}", val_acc=f"{mr_val_acc:.4f}")
pbar.update(1)
# end epoch: validate + save
tokens_trained = self._calculate_training_tokens(epoch + 1, 0)
self._ddp_barrier()
mr_val_loss, mr_val_acc = self._validate()
self._reduce(train_correct)
self._reduce(train_total)
train_acc = (train_correct.float() / train_total.clamp_min(1).float()).item()
self.checkpointing.save_checkpoint(
epoch=epoch + 1,
step=None,
train_loss=mr_step_loss,
val_loss=mr_val_loss,
tokens_trained=tokens_trained,
val_acc=mr_val_acc,
train_acc=train_acc,
)
self._ddp_barrier()
if self.rank == 0:
Logger.log(f"Epoch {epoch + 1} Complete | Val Loss: {mr_val_loss:.4f} | Val Acc: {mr_val_acc:.4f}")
def _train_ddp(self):
self.use_ddp = True
self.rank, self.world_size = setup_ddp()
try:
self.samplers = initialize_samplers_ddp(self.splits, self.rank, self.world_size)
self._train()
except Exception:
print(f"[Rank {self.rank}] Exception occurred:")
traceback.print_exc()
finally:
cleanup_ddp()
def train(self):
if self.training_config.use_ddp:
self._train_ddp()
else:
self._train()