Download checkpointing.py from tekkmaven/flint-1.2B: direct link, hf CLI and curl.
- Browser
- Download file 9.74 kB
-
https://huggingface.co/tekkmaven/flint-1.2B/resolve/main/checkpointing.py
- Command line
-
hf download hf://tekkmaven/flint-1.2B/checkpointing.py
-
curl -L -o checkpointing.py https://huggingface.co/tekkmaven/flint-1.2B/resolve/main/checkpointing.py
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_") | |
| ]) | |