XiaoyanLi's picture
Initial upload: code, training data, tokenizers, notes
83112d8 verified
Raw
History Blame Contribute Delete
16.5 kB
"""Export a training checkpoint to HuggingFace format for evaluation.
Supports ALL 5 architectures: gpt_bert, gpt2, modernized_bert, xlstm, rtd.
Produces a self-contained directory loadable via:
AutoModelForCausalLM.from_pretrained(path, trust_remote_code=True)
Usage:
python -m scripts.03_training.hf_export.export \
--checkpoint checkpoints/exp_A/checkpoint_epoch10.pt \
--output models/hf_export/exp_A \
--tokenizer models/tokenizer
"""
import argparse
import json
import shutil
from pathlib import Path
import torch
from .model_configuration import CONFIG_REGISTRY, GPTBertConfig
from .modeling import MODEL_REGISTRY
# ═══════════════════════════════════════════════════════════════════════
# State-dict translation between training models and HF wrappers
# ═══════════════════════════════════════════════════════════════════════
def _translate_gpt_bert_state_dict(state_dict: dict) -> dict:
"""Map training GPTBertModel keys to HF GPTBertForCausalLM keys.
Training model and HF model share the same parameter names because
both use the same structure. The only difference: training model
stores lm_head.tied_weight (a reference) which we drop.
"""
new_state = {}
for k, v in state_dict.items():
# Skip tied weight reference and DWA (not exported)
if k == "lm_head.tied_weight" or k.startswith("dwa."):
continue
new_state[k] = v
return new_state
def _translate_gpt2_state_dict(state_dict: dict) -> dict:
"""Map training GPT2Model keys to HF GPT2ForCausalLM keys.
Training model structure:
word_embedding, position_embedding, drop,
layers.N.attn.{ln, qkv, out_proj, attn_dropout, resid_dropout},
layers.N.ffn.{ln, fc1, fc2, dropout},
final_norm, lm_head_bias
HF model structure is identical (GPT2ForCausalLM._init_core produces
the same module tree).
"""
return {k: v for k, v in state_dict.items()}
def _translate_modern_bert_state_dict(state_dict: dict) -> dict:
"""Map training ModernBERTModel keys to HF wrapper keys.
Training model structure:
word_embedding, embed_norm, embed_dropout,
layers.N.attn.{ln, qkv, out_proj, attn_dropout, resid_dropout, rope.*},
layers.N.ffn.{ln, fc1, fc2, geglu, dropout},
final_norm, lm_head_dense, lm_head_norm, lm_head_bias
HF model mirrors this exactly.
"""
return {k: v for k, v in state_dict.items()}
def _translate_xlstm_state_dict(state_dict: dict) -> dict:
"""Map training xLSTMModel keys to HF wrapper keys."""
return {k: v for k, v in state_dict.items()}
def _translate_rtd_state_dict(state_dict: dict) -> dict:
"""Map training RTDModel discriminator keys to HF wrapper keys.
RTD training model has both generator and discriminator. We only
export the discriminator parts. Generator keys start with 'generator.'
and RTD head keys start with 'rtd_head.'.
"""
new_state = {}
for k, v in state_dict.items():
# Skip generator and RTD classification head
if k.startswith("generator.") or k.startswith("rtd_head."):
continue
new_state[k] = v
return new_state
STATE_DICT_TRANSLATORS = {
"gpt_bert": _translate_gpt_bert_state_dict,
"gpt2": _translate_gpt2_state_dict,
"modernized_bert": _translate_modern_bert_state_dict,
"xlstm": _translate_xlstm_state_dict,
"rtd": _translate_rtd_state_dict,
}
# auto_map config for each architecture
AUTO_MAP = {
"gpt_bert": {
"AutoConfig": "model_configuration.GPTBertConfig",
"AutoModel": "modeling.GPTBertModel",
"AutoModelForCausalLM": "modeling.GPTBertForCausalLM",
"AutoModelForMaskedLM": "modeling.GPTBertForMaskedLM",
},
"gpt2": {
"AutoConfig": "model_configuration.GPT2Config",
"AutoModel": "modeling.GPT2HFModel",
"AutoModelForCausalLM": "modeling.GPT2ForCausalLM",
},
"modernized_bert": {
"AutoConfig": "model_configuration.ModernBERTConfig",
"AutoModel": "modeling.ModernBERTHFModel",
"AutoModelForCausalLM": "modeling.ModernBERTForCausalLM",
"AutoModelForMaskedLM": "modeling.ModernBERTForMaskedLM",
},
"xlstm": {
"AutoConfig": "model_configuration.XLSTMConfig",
"AutoModel": "modeling.XLSTMHFModel",
"AutoModelForCausalLM": "modeling.XLSTMForCausalLM",
},
"rtd": {
"AutoConfig": "model_configuration.RTDConfig",
"AutoModel": "modeling.RTDHFModel",
"AutoModelForCausalLM": "modeling.RTDForCausalLM",
"AutoModelForMaskedLM": "modeling.RTDForMaskedLM",
},
}
# ═══════════════════════════════════════════════════════════════════════
# Main export function
# ═══════════════════════════════════════════════════════════════════════
def export_checkpoint(
checkpoint_path: str,
output_dir: str,
tokenizer_dir: str,
arch: str = None,
) -> Path:
"""Export a training checkpoint to HuggingFace format.
Args:
checkpoint_path: path to .pt checkpoint file
output_dir: output directory for HF model
tokenizer_dir: directory with tokenizer files
arch: architecture override (auto-detected from checkpoint config if None)
Returns:
Path to the output directory
"""
output = Path(output_dir)
output.mkdir(parents=True, exist_ok=True)
# ── Load checkpoint ──
print(f"Loading checkpoint: {checkpoint_path}")
ckpt = torch.load(checkpoint_path, map_location="cpu")
state_dict = ckpt["model_state_dict"]
saved_cfg = ckpt.get("config", {})
model_cfg = saved_cfg.get("model", {})
# ── Detect architecture ──
if arch is None:
arch = model_cfg.get("arch", "gpt_bert")
print(f" Architecture: {arch}")
if arch not in CONFIG_REGISTRY:
raise ValueError(f"Unknown architecture: {arch}. Available: {list(CONFIG_REGISTRY.keys())}")
# ── Build HF config ──
ConfigClass = CONFIG_REGISTRY[arch]
if arch == "gpt_bert":
config = ConfigClass(
vocab_size=model_cfg.get("vocab_size", 8192),
hidden_size=model_cfg.get("hidden_size", 384),
num_layers=model_cfg.get("num_layers", 12),
num_heads=model_cfg.get("num_heads", 6),
intermediate_size=model_cfg.get("intermediate_size", 1280),
max_seq_len=model_cfg.get("max_position_embeddings", 512),
dropout=model_cfg.get("dropout", 0.1),
layer_norm_eps=model_cfg.get("layer_norm_eps", 1e-7),
use_rope=model_cfg.get("use_rope", False),
use_geglu=model_cfg.get("use_geglu", True),
bucket_size=model_cfg.get("position_bucket_size", 32),
z_loss_weight=model_cfg.get("z_loss_weight", 1e-4),
)
elif arch == "gpt2":
config = ConfigClass(
vocab_size=model_cfg.get("vocab_size", 8192),
hidden_size=model_cfg.get("hidden_size", 384),
num_layers=model_cfg.get("num_layers", 12),
num_heads=model_cfg.get("num_heads", 6),
intermediate_size=model_cfg.get("intermediate_size", 1536),
max_seq_len=model_cfg.get("max_position_embeddings", 512),
dropout=model_cfg.get("dropout", 0.1),
layer_norm_eps=model_cfg.get("layer_norm_eps", 1e-5),
)
elif arch == "modernized_bert":
config = ConfigClass(
vocab_size=model_cfg.get("vocab_size", 8192),
hidden_size=model_cfg.get("hidden_size", 384),
num_layers=model_cfg.get("num_layers", 12),
num_heads=model_cfg.get("num_heads", 6),
intermediate_size=model_cfg.get("intermediate_size", 1280),
max_seq_len=model_cfg.get("max_position_embeddings", 512),
dropout=model_cfg.get("dropout", 0.1),
layer_norm_eps=model_cfg.get("layer_norm_eps", 1e-5),
rope_theta=model_cfg.get("rope_theta", 10000.0),
)
elif arch == "xlstm":
config = ConfigClass(
vocab_size=model_cfg.get("vocab_size", 8192),
hidden_size=model_cfg.get("hidden_size", 384),
num_layers=model_cfg.get("num_layers", 12),
num_heads=model_cfg.get("num_heads", 6),
intermediate_size=model_cfg.get("intermediate_size", 1280),
max_seq_len=model_cfg.get("max_position_embeddings", 512),
dropout=model_cfg.get("dropout", 0.1),
layer_norm_eps=model_cfg.get("layer_norm_eps", 1e-5),
)
elif arch == "rtd":
config = ConfigClass(
vocab_size=model_cfg.get("vocab_size", 8192),
hidden_size=model_cfg.get("hidden_size", 384),
num_layers=model_cfg.get("num_layers", 12),
num_heads=model_cfg.get("num_heads", 6),
intermediate_size=model_cfg.get("intermediate_size", 1280),
max_seq_len=model_cfg.get("max_position_embeddings", 512),
dropout=model_cfg.get("dropout", 0.1),
layer_norm_eps=model_cfg.get("layer_norm_eps", 1e-5),
)
config.auto_map = AUTO_MAP[arch]
# ── Translate state dict ──
translate_fn = STATE_DICT_TRANSLATORS[arch]
hf_state_dict = translate_fn(state_dict)
# ── Instantiate HF model and load weights ──
CausalLMClass = MODEL_REGISTRY[arch][0]
model = CausalLMClass(config)
missing, unexpected = model.load_state_dict(hf_state_dict, strict=False)
# Filter expected missing keys
expected_missing = {"lm_head.tied_weight"}
real_missing = [k for k in missing if k not in expected_missing]
if real_missing:
print(f" WARNING: Missing keys: {real_missing}")
if unexpected:
print(f" WARNING: Unexpected keys: {unexpected}")
# Verify weight tying for GPT-BERT
if arch == "gpt_bert":
assert model.lm_head.tied_weight is model.word_embedding.weight, "Weight tying failed!"
print(" Weight tying verified OK")
# ── Save HF format ──
# Build save state dict (exclude tied weights)
save_state = {}
for k, v in model.state_dict().items():
if k == "lm_head.tied_weight":
continue
save_state[k] = v
try:
from safetensors.torch import save_file
save_file(save_state, output / "model.safetensors")
except ImportError:
torch.save(save_state, output / "pytorch_model.bin")
config.save_pretrained(output)
print(f" Model saved to {output}")
# ── Copy code files ──
code_dir = Path(__file__).parent
for fname in ["model_configuration.py", "modeling.py"]:
shutil.copy2(code_dir / fname, output / fname)
print(" Copied model_configuration.py, modeling.py")
# ── Copy tokenizer ──
tok_dir = Path(tokenizer_dir)
if tok_dir.exists():
for f in tok_dir.iterdir():
if f.is_file():
shutil.copy2(f, output / f.name)
print(f" Copied tokenizer from {tok_dir}")
else:
print(f" WARNING: Tokenizer dir not found: {tok_dir}")
# ── Summary ──
num_params = sum(p.numel() for p in model.parameters())
print(f"\nExport complete:")
print(f" Output: {output}")
print(f" Architecture: {arch}")
print(f" Parameters: {num_params:,}")
print(f"\nTo evaluate:")
print(f" cd evaluation-pipeline-2025")
print(f" ./eval_zero_shot_fast.sh ../{output} main causal")
return output
def export_state_dict(
state_dict: dict,
arch: str,
model_cfg: dict,
output_dir: str,
tokenizer_dir: str,
) -> Path:
"""Export a model state_dict (no checkpoint wrapper) to HuggingFace format.
This is the function called by train.py after training completes,
where we already have the model state_dict and config in memory.
Args:
state_dict: model.state_dict()
arch: architecture name (gpt_bert, gpt2, etc.)
model_cfg: model config dict from config_to_dict(cfg)["model"]
output_dir: output directory for HF model
tokenizer_dir: directory with tokenizer files
Returns:
Path to the output directory
"""
output = Path(output_dir)
output.mkdir(parents=True, exist_ok=True)
print(f" Exporting HF model ({arch}) to {output}")
if arch not in CONFIG_REGISTRY:
raise ValueError(f"Unknown architecture: {arch}. Available: {list(CONFIG_REGISTRY.keys())}")
# ── Build HF config ──
ConfigClass = CONFIG_REGISTRY[arch]
# Generic config construction
config_kwargs = {
"vocab_size": model_cfg.get("vocab_size", 8192),
"hidden_size": model_cfg.get("hidden_size", 384),
"num_layers": model_cfg.get("num_layers", 12),
"num_heads": model_cfg.get("num_heads", 6),
"intermediate_size": model_cfg.get("intermediate_size", 1280),
"max_seq_len": model_cfg.get("max_position_embeddings", 512),
"dropout": model_cfg.get("dropout", 0.1),
"layer_norm_eps": model_cfg.get("layer_norm_eps", 1e-7),
}
# Arch-specific extras
if arch == "gpt_bert":
config_kwargs.update({
"use_rope": model_cfg.get("use_rope", False),
"use_geglu": model_cfg.get("use_geglu", True),
"bucket_size": model_cfg.get("position_bucket_size", 32),
"z_loss_weight": model_cfg.get("z_loss_weight", 1e-4),
})
elif arch == "modernized_bert":
config_kwargs.update({
"rope_theta": model_cfg.get("rope_theta", 10000.0),
})
config = ConfigClass(**config_kwargs)
config.auto_map = AUTO_MAP[arch]
# ── Translate state dict ──
translate_fn = STATE_DICT_TRANSLATORS[arch]
hf_state_dict = translate_fn(state_dict)
# ── Instantiate HF model and load weights ──
CausalLMClass = MODEL_REGISTRY[arch][0]
model = CausalLMClass(config)
missing, unexpected = model.load_state_dict(hf_state_dict, strict=False)
expected_missing = {"lm_head.tied_weight"}
real_missing = [k for k in missing if k not in expected_missing]
if real_missing:
print(f" WARNING: Missing keys in HF model: {real_missing}")
if unexpected:
print(f" WARNING: Unexpected keys in HF model: {unexpected}")
# ── Save HF format ──
save_state = {k: v for k, v in model.state_dict().items()
if k != "lm_head.tied_weight"}
try:
from safetensors.torch import save_file
save_file(save_state, output / "model.safetensors")
except ImportError:
torch.save(save_state, output / "pytorch_model.bin")
config.save_pretrained(output)
# ── Copy code files ──
code_dir = Path(__file__).parent
for fname in ["model_configuration.py", "modeling.py"]:
shutil.copy2(code_dir / fname, output / fname)
# ── Copy tokenizer ──
tok_dir = Path(tokenizer_dir)
if tok_dir.exists():
for f in tok_dir.iterdir():
if f.is_file():
shutil.copy2(f, output / f.name)
print(f" Tokenizer copied from {tok_dir}")
else:
print(f" WARNING: Tokenizer dir not found: {tok_dir}")
num_params = sum(p.numel() for p in model.parameters())
print(f" HF export complete: {num_params:,} params -> {output}")
return output
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Export BabyLM checkpoint to HF format")
parser.add_argument("--checkpoint", required=True, help="Path to training checkpoint .pt")
parser.add_argument("--output", required=True, help="Output directory for HF model")
parser.add_argument("--tokenizer", default="models/tokenizer", help="Tokenizer directory")
parser.add_argument("--arch", default=None, help="Architecture override (auto-detected if omitted)")
args = parser.parse_args()
export_checkpoint(args.checkpoint, args.output, args.tokenizer, args.arch)