flint-1.2B / checkpointing.py
tekkmaven's picture
Fix: push full model weights to Hub at hub_push_interval (not just status JSON)"
5a97dba verified
Raw History Blame Contribute Delete
9.74 kB
"""
Flint-1.2B Checkpoint Manager
==============================
- Keeps only 2 local checkpoints (Kaggle 20GB disk limit)
- Pushes full model weights to HF Hub periodically
- Saves as float16 numpy (bfloat16 not supported by numpy natively)
"""
import os
import json
import time
import shutil
from pathlib import Path
from typing import Optional, Dict, Any, Tuple
import numpy as np
import jax
import jax.numpy as jnp
class CheckpointManager:
"""Checkpoint manager: local saves + Hub weight pushes."""
MAX_CHECKPOINTS = 2
def __init__(self, config):
self.config = config
self.ckpt_dir = Path(config.checkpoint_dir)
self.ckpt_dir.mkdir(parents=True, exist_ok=True)
self.hub_model_id = config.hub_model_id
def save(self, step, params, opt_state, rng_key, data_position, metrics):
"""Save locally + push weights to Hub at intervals."""
self._rotate_before_save()
t0 = time.time()
step_dir = self.ckpt_dir / f"step_{step:07d}"
step_dir.mkdir(parents=True, exist_ok=True)
# Save params as float16
param_leaves = jax.tree_util.tree_leaves(params)
param_dir = step_dir / "params"
param_dir.mkdir(exist_ok=True)
for i, p in enumerate(param_leaves):
arr = np.array(p, dtype=np.float32).astype(np.float16)
np.save(str(param_dir / f"p_{i:04d}.npy"), arr)
# RNG key
rng_arr = np.array(rng_key, dtype=np.uint32)
np.save(str(step_dir / "rng_key.npy"), rng_arr)
# Metadata
meta = {
'step': step,
'data_position': int(data_position),
'metrics': {k: float(v) if hasattr(v, '__float__') else str(v) for k, v in metrics.items()},
'timestamp': time.strftime('%Y-%m-%d %H:%M:%S'),
'num_param_files': len(param_leaves),
'dtype': 'float16',
}
with open(step_dir / "metadata.json", 'w') as f:
json.dump(meta, f, indent=2)
elapsed = time.time() - t0
size_mb = sum(f.stat().st_size for f in step_dir.rglob("*") if f.is_file()) / 1e6
print(f"[Checkpoint] Saved step {step} ({size_mb:.0f}MB) in {elapsed:.1f}s")
# Push FULL WEIGHTS to Hub at hub_push_interval
if step % self.config.hub_push_interval == 0:
self._push_weights_to_hub(step, param_leaves, meta)
def _push_weights_to_hub(self, step, param_leaves, meta):
"""Push full model weights to HF Hub as a compressed numpy archive."""
try:
from huggingface_hub import HfApi
api = HfApi()
t0 = time.time()
print(f"[Hub] Pushing full weights at step {step}...")
# Save all params as a single compressed .npz file (~1.1GB for fp16)
staging_dir = self.ckpt_dir / "_hub_staging"
staging_dir.mkdir(exist_ok=True)
# Create the weight file
weight_dict = {}
for i, p in enumerate(param_leaves):
weight_dict[f"p_{i:04d}"] = np.array(p, dtype=np.float32).astype(np.float16)
weights_path = str(staging_dir / "model_weights.npz")
np.savez_compressed(weights_path, **weight_dict)
weight_size_mb = os.path.getsize(weights_path) / 1e6
# Upload weights
api.upload_file(
path_or_fileobj=weights_path,
path_in_repo=f"checkpoints/step_{step:07d}_weights.npz",
repo_id=self.hub_model_id,
commit_message=f"Checkpoint weights: step {step} ({weight_size_mb:.0f}MB)",
)
# Upload metadata
meta_path = str(staging_dir / "training_status.json")
with open(meta_path, 'w') as f:
json.dump(meta, f, indent=2)
api.upload_file(
path_or_fileobj=meta_path,
path_in_repo="training_status.json",
repo_id=self.hub_model_id,
commit_message=f"Training status: step {step}",
)
elapsed = time.time() - t0
print(f"[Hub] ✓ Pushed {weight_size_mb:.0f}MB weights in {elapsed:.0f}s → "
f"checkpoints/step_{step:07d}_weights.npz")
# Cleanup staging
shutil.rmtree(staging_dir, ignore_errors=True)
except Exception as e:
print(f"[Hub] Push failed (non-fatal): {e}")
# Clean staging dir even on failure
staging_dir = self.ckpt_dir / "_hub_staging"
if staging_dir.exists():
shutil.rmtree(staging_dir, ignore_errors=True)
def _rotate_before_save(self):
"""Delete oldest checkpoint BEFORE saving."""
step_dirs = sorted([
d for d in self.ckpt_dir.iterdir()
if d.is_dir() and d.name.startswith("step_")
])
while len(step_dirs) >= self.MAX_CHECKPOINTS:
old = step_dirs.pop(0)
print(f"[Checkpoint] Deleting old: {old.name} (freeing space)")
shutil.rmtree(old, ignore_errors=True)
def load_latest(self) -> Tuple[Optional[Dict], int]:
"""Load most recent checkpoint."""
step_dirs = sorted([
d for d in self.ckpt_dir.iterdir()
if d.is_dir() and d.name.startswith("step_")
])
if not step_dirs:
return None, 0
for step_dir in reversed(step_dirs):
step = int(step_dir.name.split("_")[1])
result, loaded_step = self._load_from_dir(step_dir, step)
if result is not None:
return result, loaded_step
return None, 0
def load_specific(self, step) -> Tuple[Optional[Dict], int]:
"""Load a specific step."""
step_dir = self.ckpt_dir / f"step_{step:07d}"
if not step_dir.exists():
print(f"[Checkpoint] Step {step} not found")
return None, 0
return self._load_from_dir(step_dir, step)
def _load_from_dir(self, step_dir, step) -> Tuple[Optional[Dict], int]:
"""Load checkpoint. Converts float16 → bfloat16."""
try:
param_dir = step_dir / "params"
if not param_dir.exists():
return None, 0
param_files = sorted(param_dir.glob("p_*.npy"))
if not param_files:
return None, 0
param_leaves = []
for f in param_files:
arr = np.load(str(f))
if arr.dtype == np.float16 or arr.dtype == np.float32:
param_leaves.append(jnp.array(arr, dtype=jnp.bfloat16))
elif arr.dtype.kind == 'V':
raw = arr.view(np.uint16).view(np.float16)
param_leaves.append(jnp.array(raw, dtype=jnp.bfloat16))
else:
param_leaves.append(jnp.array(arr, dtype=jnp.bfloat16))
rng_key = None
rng_path = step_dir / "rng_key.npy"
if rng_path.exists():
rng_arr = np.load(str(rng_path))
if rng_arr.dtype.kind != 'V':
rng_key = jnp.array(rng_arr, dtype=jnp.uint32)
meta = {}
meta_path = step_dir / "metadata.json"
if meta_path.exists():
with open(meta_path) as f:
meta = json.load(f)
state = {
'params_leaves': param_leaves,
'rng_key': rng_key,
'data_position': meta.get('data_position', 0),
}
print(f"[Checkpoint] Loaded step {step} ({len(param_leaves)} tensors)")
return state, step
except Exception as e:
print(f"[Checkpoint] Failed to load {step_dir}: {e}")
return None, 0
def load_from_hub(self, step=None):
"""Download weights from Hub (for cross-session resume)."""
try:
from huggingface_hub import HfApi, hf_hub_download
api = HfApi()
# List available weight files
files = api.list_repo_files(self.hub_model_id)
weight_files = sorted([f for f in files if f.startswith("checkpoints/") and f.endswith("_weights.npz")])
if not weight_files:
print("[Hub] No weight files found on Hub")
return None, 0
# Pick specific step or latest
if step:
target = f"checkpoints/step_{step:07d}_weights.npz"
if target not in weight_files:
print(f"[Hub] Step {step} not found on Hub. Available: {weight_files}")
return None, 0
dl_file = target
else:
dl_file = weight_files[-1]
# Extract step from filename
loaded_step = int(dl_file.split("step_")[1].split("_")[0])
print(f"[Hub] Downloading weights from step {loaded_step}...")
local_path = hf_hub_download(self.hub_model_id, dl_file)
data = np.load(local_path)
param_leaves = []
for key in sorted(data.files):
param_leaves.append(jnp.array(data[key], dtype=jnp.bfloat16))
print(f"[Hub] ✓ Loaded {len(param_leaves)} tensors from step {loaded_step}")
return {'params_leaves': param_leaves, 'rng_key': None, 'data_position': 0}, loaded_step
except Exception as e:
print(f"[Hub] Download failed: {e}")
return None, 0
def list_checkpoints(self):
"""List local checkpoint steps."""
return sorted([
int(d.name.split("_")[1])
for d in self.ckpt_dir.iterdir()
if d.is_dir() and d.name.startswith("step_")
])