""" Modal training script for Retriever500M. Runs pretraining and/or SFT on Modal cloud GPUs instead of the local laptop. Supports resuming from the existing bf16 checkpoint. Usage: # Step 1: Seed the volume with the initial checkpoint modal run src/train_modal.py --command seed # Step 2: Run pretraining (continues from latest checkpoint) modal run src/train_modal.py --command pretrain --steps 1000 # Step 3: Run SFT (starts from pretrained checkpoint) modal run src/train_modal.py --command sft --steps 500 # Step 4: Download checkpoints back to local modal run src/train_modal.py --command download # Or do everything in one shot: modal run src/train_modal.py --command pipeline --pretrain-steps 1000 --sft-steps 500 """ import os import sys import json import time import argparse from dataclasses import asdict import numpy as np import torch import torch.nn.functional as F from tqdm import tqdm import modal # ─── Modal Setup ───────────────────────────────────────────────────────────── APP_NAME = "retriever500m" VOLUME_NAME = "retriever500m-data" PROJECT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) SRC_DIR = os.path.join(PROJECT_DIR, "src") DATA_DIR = os.path.join(PROJECT_DIR, "data") TOKENIZER_DIR = os.path.join(PROJECT_DIR, "tokenizer") CHECKPOINT_DIR = os.path.join(PROJECT_DIR, "checkpoints") volume = modal.Volume.from_name(VOLUME_NAME, create_if_missing=True) # Image: PyTorch with CUDA, plus bitsandbytes for 8-bit optimizer. # Data, tokenizer, src, and initial checkpoint are baked into the image. # The volume is used only for saving new checkpoints and logs. image = ( modal.Image.from_registry( "pytorch/pytorch:2.4.1-cuda12.1-cudnn9-runtime", add_python="3.11", ) .pip_install("tokenizers>=0.15,<0.21", "tqdm", "numpy") .pip_install("bitsandbytes>=0.43,<0.45") .add_local_dir(SRC_DIR, "/root/src") .add_local_dir(TOKENIZER_DIR, "/root/tokenizer") .add_local_dir(DATA_DIR, "/root/data", ignore=["raw_large/", "raw/", "dedup/"]) .add_local_dir(CHECKPOINT_DIR, "/root/seed_checkpoints") ) app = modal.App(APP_NAME, image=image) # Remote paths inside the container. # Data/tokenizer/src are baked into the image. # The volume is mounted at /root/vol for checkpoints and logs. VOL_MOUNT = "/root/vol" REMOTE_DATA = "/root/data" # baked into image REMOTE_CKPT = "/root/vol/checkpoints" # on volume (for new checkpoints) REMOTE_LOGS = "/root/vol/logs" # on volume REMOTE_TOKENIZER = "/root/tokenizer" # baked into image REMOTE_SRC = "/root/src" # baked into image REMOTE_SEED_CKPT = "/root/seed_checkpoints" # baked into image (initial checkpoint) # ─── Seed command ──────────────────────────────────────────────────────────── @app.function(volumes={VOL_MOUNT: volume}) def seed(): """Copy the initial checkpoint from the image to the volume.""" import shutil os.makedirs(REMOTE_CKPT, exist_ok=True) os.makedirs(REMOTE_LOGS, exist_ok=True) # Copy seed checkpoints from image to volume if os.path.exists(REMOTE_SEED_CKPT): for fname in os.listdir(REMOTE_SEED_CKPT): src = os.path.join(REMOTE_SEED_CKPT, fname) dst = os.path.join(REMOTE_CKPT, fname) if os.path.isfile(src) and not os.path.exists(dst): shutil.copy2(src, dst) print(f" Seeded {fname} ({os.path.getsize(src) / 1e6:.1f} MB)") elif os.path.exists(dst): print(f" SKIP {fname} (already on volume)") volume.commit() print("Seed complete.") # ─── Download command ──────────────────────────────────────────────────────── @app.function(volumes={VOL_MOUNT: volume}) def download(): """Download checkpoints and logs from the Modal volume to local.""" import shutil os.makedirs(CHECKPOINT_DIR, exist_ok=True) os.makedirs(os.path.join(PROJECT_DIR, "logs"), exist_ok=True) # Download checkpoints for fname in os.listdir(REMOTE_CKPT): src = os.path.join(REMOTE_CKPT, fname) if os.path.isfile(src): dst = os.path.join(CHECKPOINT_DIR, fname) shutil.copy2(src, dst) print(f" Downloaded {fname} ({os.path.getsize(src) / 1e6:.1f} MB)") # Download logs for fname in os.listdir(REMOTE_LOGS): src = os.path.join(REMOTE_LOGS, fname) if os.path.isfile(src): dst = os.path.join(PROJECT_DIR, "logs", fname) shutil.copy2(src, dst) print(f" Downloaded log {fname}") print("Download complete.") # ─── Training logic (shared) ───────────────────────────────────────────────── def get_lr(step, warmup, max_steps, max_lr, min_lr): """Cosine LR schedule with linear warmup.""" if step < warmup: return max_lr * (step + 1) / warmup if step > max_steps: return min_lr decay_ratio = (step - warmup) / (max_steps - warmup) coeff = 0.5 * (1.0 + np.cos(np.pi * decay_ratio)) return min_lr + coeff * (max_lr - min_lr) def setup_optimizer(model, lr, use_8bit=True): decay_params, no_decay_params = [], [] for name, param in model.named_parameters(): if not param.requires_grad: continue if "embedding" in name or "norm" in name: no_decay_params.append(param) else: decay_params.append(param) param_groups = [ {"params": decay_params, "weight_decay": 0.1}, {"params": no_decay_params, "weight_decay": 0.0}, ] if use_8bit: try: import bitsandbytes as bnb optimizer = bnb.optim.AdamW8bit(param_groups, lr=lr, betas=(0.9, 0.95), eps=1e-8) print("Using 8-bit AdamW (bitsandbytes)") return optimizer except Exception as e: print(f"8-bit optimizer unavailable ({e}), falling back to AdamW") optimizer = torch.optim.AdamW(param_groups, lr=lr, betas=(0.9, 0.95), eps=1e-8) print("Using standard AdamW") return optimizer # ─── Pretraining ───────────────────────────────────────────────────────────── def load_and_tokenize(corpus_path, tokenizer): """Load corpus, tokenize, return flat numpy array of token IDs.""" print(f"Loading corpus from {corpus_path}...") with open(corpus_path, "r", encoding="utf-8") as f: text = f.read() print(f"Corpus size: {len(text) / 1e6:.1f} MB") chunk_size = 1_000_000 all_tokens = [] print("Tokenizing corpus...") for i in tqdm(range(0, len(text), chunk_size)): chunk = text[i : i + chunk_size] encoded = tokenizer.encode(chunk) all_tokens.extend(encoded.ids) tokens = np.array(all_tokens, dtype=np.int32) print(f"Total tokens: {len(tokens):,}") return tokens def get_batch(tokens, batch_size, seq_len, device): max_start = len(tokens) - seq_len - 1 indices = np.random.randint(0, max_start, size=batch_size) input_ids = np.stack([tokens[i : i + seq_len] for i in indices]) targets = np.stack([tokens[i + 1 : i + seq_len + 1] for i in indices]) input_ids = torch.from_numpy(input_ids).long().to(device) targets = torch.from_numpy(targets).long().to(device) return input_ids, targets def run_pretrain(steps, batch_size, grad_accum, seq_len, lr, warmup, save_every, log_every, corpus, resume, use_8bit_adam): sys.path.insert(0, REMOTE_SRC) from model import ModelConfig, Retriever500M from tokenizers import Tokenizer device = torch.device("cuda") print(f"Device: {device}") print(f"GPU: {torch.cuda.get_device_name(0)}") print(f"VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB") os.makedirs(REMOTE_CKPT, exist_ok=True) os.makedirs(REMOTE_LOGS, exist_ok=True) # Tokenizer tokenizer_path = os.path.join(REMOTE_TOKENIZER, "tokenizer.json") tokenizer = Tokenizer.from_file(tokenizer_path) vocab_size = tokenizer.get_vocab_size() print(f"Vocab size: {vocab_size}") # Data if corpus == "curated": corpus_path = os.path.join(REMOTE_DATA, "corpus_curated.txt") elif corpus == "default": corpus_path = os.path.join(REMOTE_DATA, "corpus.txt") else: corpus_path = corpus tokens = load_and_tokenize(corpus_path, tokenizer) # Model config = ModelConfig( vocab_size=vocab_size, d_model=1_280, n_layers=23, n_heads=20, d_ff=3_456, max_seq_len=seq_len, dropout=0.0, tie_embeddings=True, ) model = Retriever500M(config).to(device) total_params = model.count_parameters() print(f"Model parameters: {total_params:,} ({total_params / 1e6:.1f}M)") optimizer = setup_optimizer(model, lr, use_8bit=use_8bit_adam) # Resume start_step = 0 best_loss = float("inf") accum_loss = 0.0 prev_log_steps = [] if resume: resume_path = os.path.join(REMOTE_CKPT, "latest.pt") if not os.path.exists(resume_path): resume_path = os.path.join(REMOTE_CKPT, "latest_bf16.pt") if not os.path.exists(resume_path): # Fall back to seed checkpoint in image resume_path = os.path.join(REMOTE_SEED_CKPT, "latest_bf16.pt") if os.path.exists(resume_path): print(f"Resuming from {resume_path}") ckpt = torch.load(resume_path, map_location=device, weights_only=False) model.load_state_dict(ckpt["model_state_dict"]) start_step = int(ckpt.get("step", 0)) best_loss = float(ckpt.get("loss", float("inf"))) accum_loss = best_loss print(f" Resuming at step {start_step} (best_loss={best_loss:.4f})") prev_log_path = os.path.join(REMOTE_LOGS, "training_log.json") if os.path.exists(prev_log_path): try: with open(prev_log_path, "r") as f: prev_log = json.load(f) prev_log_steps = prev_log.get("steps", []) print(f" Loaded {len(prev_log_steps)} previous log entries") except Exception: pass else: print("WARNING: No checkpoint found, starting from scratch!") # Training loop effective_batch = batch_size * grad_accum max_steps_total = start_step + steps print(f"\nTraining configuration:") print(f" Micro batch size: {batch_size}") print(f" Gradient accum: {grad_accum}") print(f" Effective batch: {effective_batch}") print(f" Sequence length: {seq_len}") print(f" Learning rate: {lr}") print(f" Steps this run: {steps}") print(f" Start step: {start_step}") print(f" Target step: {max_steps_total}") print() log = { "config": asdict(config), "train_args": {"steps": steps, "batch_size": batch_size, "grad_accum": grad_accum, "seq_len": seq_len, "lr": lr, "warmup": warmup, "save_every": save_every, "log_every": log_every, "corpus": corpus, "resume": resume, "use_8bit_adam": use_8bit_adam}, "total_params": total_params, "steps": list(prev_log_steps), } model.train() start_time = time.time() pbar = tqdm(range(start_step, max_steps_total), desc="Training", initial=start_step, total=max_steps_total) for step in pbar: lr_now = get_lr(step, warmup, max_steps_total, lr, lr * 0.1) for pg in optimizer.param_groups: pg["lr"] = lr_now optimizer.zero_grad(set_to_none=True) total_loss = 0.0 for _ in range(grad_accum): input_ids, targets = get_batch(tokens, batch_size, seq_len, device) with torch.autocast(device_type="cuda", dtype=torch.bfloat16): out = model(input_ids, targets=targets, use_checkpoint=False) loss = out["loss"] / grad_accum loss.backward() total_loss += loss.item() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() avg_loss = total_loss accum_loss = accum_loss * 0.95 + avg_loss * 0.05 if step % log_every == 0 or step == max_steps_total - 1: elapsed = time.time() - start_time steps_this_run = step - start_step + 1 steps_per_sec = steps_this_run / elapsed vram_used = torch.cuda.max_memory_allocated() / 1e9 log_entry = { "step": step, "loss": avg_loss, "ema_loss": accum_loss, "lr": lr_now, "elapsed_s": elapsed, "steps_per_sec": steps_per_sec, "vram_gb": vram_used, } log["steps"].append(log_entry) pbar.set_postfix({ "loss": f"{avg_loss:.4f}", "ema": f"{accum_loss:.4f}", "lr": f"{lr_now:.2e}", "vram": f"{vram_used:.1f}G", }) if (step + 1) % save_every == 0 or step == max_steps_total - 1: ckpt_path = os.path.join(REMOTE_CKPT, f"model_step_{step + 1}.pt") torch.save({ "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "config": asdict(config), "step": step + 1, "loss": accum_loss, }, ckpt_path) print(f"\n Saved checkpoint: {ckpt_path}") latest_path = os.path.join(REMOTE_CKPT, "latest.pt") torch.save({ "model_state_dict": model.state_dict(), "config": asdict(config), "step": step + 1, "loss": accum_loss, }, latest_path) if accum_loss < best_loss: best_loss = accum_loss best_path = os.path.join(REMOTE_CKPT, "best.pt") torch.save({ "model_state_dict": model.state_dict(), "config": asdict(config), "step": step + 1, "loss": accum_loss, }, best_path) # Save log log_path = os.path.join(REMOTE_LOGS, "training_log.json") with open(log_path, "w") as f: json.dump(log, f, indent=2) # Commit volume so checkpoints persist volume.commit() if step % 50 == 0: torch.cuda.reset_peak_memory_stats() # Final log save log_path = os.path.join(REMOTE_LOGS, "training_log.json") with open(log_path, "w") as f: json.dump(log, f, indent=2) volume.commit() total_time = time.time() - start_time print(f"\nPretraining complete!") print(f" Total time: {total_time:.1f}s ({total_time/60:.1f} min)") print(f" Final EMA loss: {accum_loss:.4f}") print(f" Best loss: {best_loss:.4f}") return accum_loss # ─── SFT ───────────────────────────────────────────────────────────────────── # Special token IDs SYSTEM_ID = 32000 USER_ID = 32001 ASSISTANT_ID = 32002 SEARCH_ID = 32003 RESULT_ID = 32004 EVIDENCE_ID = 32005 REASONING_ID = 32006 FINISH_ID = 32007 END_ID = 32008 def load_special_tokens(): global SYSTEM_ID, USER_ID, ASSISTANT_ID, SEARCH_ID, RESULT_ID global EVIDENCE_ID, REASONING_ID, FINISH_ID, END_ID path = os.path.join(REMOTE_TOKENIZER, "special_tokens.json") if os.path.exists(path): with open(path, "r") as f: data = json.load(f) ids = data["token_ids"] SYSTEM_ID = ids.get("", 32000) USER_ID = ids.get("", 32001) ASSISTANT_ID = ids.get("", 32002) SEARCH_ID = ids.get("<|search|>", 32003) RESULT_ID = ids.get("<|result|>", 32004) EVIDENCE_ID = ids.get("<|evidence|>", 32005) REASONING_ID = ids.get("<|reasoning|>", 32006) FINISH_ID = ids.get("<|finish|>", 32007) END_ID = ids.get("<|end|>", 32008) def format_trace_to_tokens(trace, tokenizer, max_seq_len=768): messages = trace["trace"] all_tokens = [] loss_mask = [] for msg in messages: role = msg["role"] content = msg["content"] if role == "system": tokens = [SYSTEM_ID] + tokenizer.encode(content).ids + [END_ID] all_tokens.extend(tokens) loss_mask.extend([0] * len(tokens)) elif role == "user": tokens = [USER_ID] + tokenizer.encode(content).ids + [END_ID] all_tokens.extend(tokens) loss_mask.extend([0] * len(tokens)) elif role == "assistant": tokens = [ASSISTANT_ID] + tokenizer.encode(content).ids + [END_ID] all_tokens.extend(tokens) loss_mask.extend([1] * len(tokens)) elif role == "result": if content: tokens = [RESULT_ID] + tokenizer.encode(content).ids + [END_ID] else: tokens = [RESULT_ID, END_ID] all_tokens.extend(tokens) loss_mask.extend([0] * len(tokens)) if len(all_tokens) > max_seq_len: all_tokens = all_tokens[:max_seq_len] loss_mask = loss_mask[:max_seq_len] return np.array(all_tokens, dtype=np.int32), np.array(loss_mask, dtype=np.int32) def load_sft_dataset(traces_path, tokenizer, max_seq_len=768): print(f"Loading SFT traces from {traces_path}...") dataset = [] with open(traces_path, "r", encoding="utf-8") as f: for line in f: trace = json.loads(line) ids, mask = format_trace_to_tokens(trace, tokenizer, max_seq_len) if len(ids) > 10: dataset.append((ids, mask)) print(f" Loaded {len(dataset):,} traces") return dataset def get_sft_batch(dataset, batch_size, seq_len, device): indices = np.random.randint(0, len(dataset), size=batch_size) input_ids_list = [] loss_mask_list = [] for idx in indices: ids, mask = dataset[idx] if len(ids) < seq_len: pad_len = seq_len - len(ids) ids = np.concatenate([ids, np.zeros(pad_len, dtype=np.int32)]) mask = np.concatenate([mask, np.zeros(pad_len, dtype=np.int32)]) else: ids = ids[:seq_len] mask = mask[:seq_len] input_ids_list.append(ids) loss_mask_list.append(mask) input_ids = torch.from_numpy(np.stack(input_ids_list)).long().to(device) loss_mask = torch.from_numpy(np.stack(loss_mask_list)).long().to(device) targets = torch.cat([input_ids[:, 1:], torch.zeros_like(input_ids[:, :1])], dim=1) return input_ids, targets, loss_mask def run_sft(steps, batch_size, grad_accum, seq_len, lr, warmup, save_every, log_every, use_8bit_adam): from torch import nn sys.path.insert(0, REMOTE_SRC) from model import ModelConfig, Retriever500M from tokenizers import Tokenizer load_special_tokens() device = torch.device("cuda") print(f"Device: {device}") print(f"GPU: {torch.cuda.get_device_name(0)}") print(f"VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB") os.makedirs(REMOTE_CKPT, exist_ok=True) os.makedirs(REMOTE_LOGS, exist_ok=True) # Tokenizer (agent tokenizer with special tokens) tokenizer_path = os.path.join(REMOTE_TOKENIZER, "tokenizer_agent.json") tokenizer = Tokenizer.from_file(tokenizer_path) vocab_size = tokenizer.get_vocab_size() print(f"Vocab size: {vocab_size}") # Data traces_path = os.path.join(REMOTE_DATA, "sft_traces.jsonl") gold_path = os.path.join(REMOTE_DATA, "gold_traces.jsonl") dataset = load_sft_dataset(traces_path, tokenizer, seq_len) gold_dataset = load_sft_dataset(gold_path, tokenizer, seq_len) dataset.extend(gold_dataset) print(f" Total (with gold): {len(dataset):,}") # Model config = ModelConfig( vocab_size=vocab_size, d_model=1_280, n_layers=23, n_heads=20, d_ff=3_456, max_seq_len=seq_len, dropout=0.0, tie_embeddings=True, ) model = Retriever500M(config).to(device) # Load pretrained checkpoint ckpt_path = os.path.join(REMOTE_CKPT, "latest.pt") if not os.path.exists(ckpt_path): ckpt_path = os.path.join(REMOTE_CKPT, "latest_bf16.pt") if not os.path.exists(ckpt_path): ckpt_path = os.path.join(REMOTE_SEED_CKPT, "latest_bf16.pt") if os.path.exists(ckpt_path): print(f"Loading pretrained weights from {ckpt_path}...") ckpt = torch.load(ckpt_path, map_location=device, weights_only=False) old_config = ModelConfig(**ckpt["config"]) state_dict = ckpt["model_state_dict"] old_vocab = old_config.vocab_size if old_vocab != vocab_size: print(f" Vocab size mismatch: {old_vocab} -> {vocab_size}") old_weight = state_dict["token_embedding.weight"] d_model = old_weight.shape[1] new_weight = torch.zeros(vocab_size, d_model) new_weight[:old_vocab] = old_weight nn.init.normal_(new_weight[old_vocab:], mean=0.0, std=0.02) state_dict["token_embedding.weight"] = new_weight model.load_state_dict(state_dict) print(f" Loaded (step {ckpt.get('step', '?')}, loss {ckpt.get('loss', '?')})") else: print("WARNING: No checkpoint found, starting from scratch!") total_params = model.count_parameters() print(f"Model parameters: {total_params:,} ({total_params / 1e6:.1f}M)") optimizer = setup_optimizer(model, lr, use_8bit=use_8bit_adam) effective_batch = batch_size * grad_accum print(f"\nSFT configuration:") print(f" Batch size: {batch_size}") print(f" Grad accum: {grad_accum}") print(f" Effective batch: {effective_batch}") print(f" Sequence length: {seq_len}") print(f" Learning rate: {lr}") print(f" Steps: {steps}") print(f" Warmup: {warmup}") print() log = { "config": asdict(config), "train_args": {"steps": steps, "batch_size": batch_size, "grad_accum": grad_accum, "seq_len": seq_len, "lr": lr, "warmup": warmup, "save_every": save_every, "log_every": log_every, "use_8bit_adam": use_8bit_adam}, "total_params": total_params, "steps": [], } model.train() start_time = time.time() accum_loss = 0.0 best_loss = float("inf") pbar = tqdm(range(steps), desc="SFT") for step in pbar: lr_now = get_lr(step, warmup, steps, lr, lr * 0.1) for pg in optimizer.param_groups: pg["lr"] = lr_now optimizer.zero_grad(set_to_none=True) total_loss = 0.0 for _ in range(grad_accum): input_ids, targets, loss_mask = get_sft_batch( dataset, batch_size, seq_len, device ) with torch.autocast(device_type="cuda", dtype=torch.bfloat16): out = model(input_ids, targets=targets, use_checkpoint=True) logits = out["logits"] if loss_mask.sum() > 0: shifted_mask = loss_mask[:, 1:].contiguous() masked_logits = logits[:, :-1, :].contiguous() masked_targets = targets[:, :-1].contiguous() flat_logits = masked_logits.view(-1, masked_logits.size(-1)) flat_targets = masked_targets.view(-1) flat_mask = shifted_mask.view(-1).float() per_token_loss = F.cross_entropy( flat_logits, flat_targets, ignore_index=-100, reduction="none" ) masked_loss = (per_token_loss * flat_mask).sum() / flat_mask.sum().clamp(min=1) loss = masked_loss / grad_accum else: loss = out["loss"] / grad_accum loss.backward() total_loss += loss.item() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() avg_loss = total_loss accum_loss = accum_loss * 0.95 + avg_loss * 0.05 if step % log_every == 0 or step == steps - 1: elapsed = time.time() - start_time steps_per_sec = (step + 1) / elapsed vram_used = torch.cuda.max_memory_allocated() / 1e9 log_entry = { "step": step, "loss": avg_loss, "ema_loss": accum_loss, "lr": lr_now, "elapsed_s": elapsed, "steps_per_sec": steps_per_sec, "vram_gb": vram_used, } log["steps"].append(log_entry) pbar.set_postfix({ "loss": f"{avg_loss:.4f}", "ema": f"{accum_loss:.4f}", "lr": f"{lr_now:.2e}", "vram": f"{vram_used:.1f}G", }) if (step + 1) % save_every == 0 or step == steps - 1: ckpt_path = os.path.join(REMOTE_CKPT, f"sft_step_{step + 1}.pt") torch.save({ "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "config": asdict(config), "step": step + 1, "loss": accum_loss, }, ckpt_path) print(f"\n Saved checkpoint: {ckpt_path}") latest_path = os.path.join(REMOTE_CKPT, "sft_latest.pt") torch.save({ "model_state_dict": model.state_dict(), "config": asdict(config), "step": step + 1, "loss": accum_loss, }, latest_path) if accum_loss < best_loss: best_loss = accum_loss best_path = os.path.join(REMOTE_CKPT, "sft_best.pt") torch.save({ "model_state_dict": model.state_dict(), "config": asdict(config), "step": step + 1, "loss": accum_loss, }, best_path) log_path = os.path.join(REMOTE_LOGS, "sft_log.json") with open(log_path, "w") as f: json.dump(log, f, indent=2) volume.commit() if step % 50 == 0: torch.cuda.reset_peak_memory_stats() log_path = os.path.join(REMOTE_LOGS, "sft_log.json") with open(log_path, "w") as f: json.dump(log, f, indent=2) volume.commit() total_time = time.time() - start_time print(f"\nSFT complete!") print(f" Total time: {total_time:.1f}s ({total_time/60:.1f} min)") print(f" Final EMA loss: {accum_loss:.4f}") print(f" Best loss: {best_loss:.4f}") return accum_loss # ─── Modal entry points ────────────────────────────────────────────────────── GPU_CHOICES = {"a10g": "A10G", "a100": "A100", "h100": "H100"} @app.function( volumes={VOL_MOUNT: volume}, gpu="A10G", timeout=3600, ) def pretrain( steps=1000, batch_size=8, grad_accum=4, seq_len=512, lr=3e-4, warmup=100, save_every=200, log_every=10, corpus="curated", resume=True, use_8bit_adam=True, gpu="a10g", ): """Run pretraining on Modal GPU.""" return run_pretrain(steps, batch_size, grad_accum, seq_len, lr, warmup, save_every, log_every, corpus, resume, use_8bit_adam) @app.function( volumes={VOL_MOUNT: volume}, gpu="A10G", timeout=3600, env={"PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True"}, ) def sft( steps=500, batch_size=4, grad_accum=8, seq_len=768, lr=5e-5, warmup=20, save_every=100, log_every=10, use_8bit_adam=True, gpu="a10g", ): """Run SFT on Modal GPU.""" return run_sft(steps, batch_size, grad_accum, seq_len, lr, warmup, save_every, log_every, use_8bit_adam) @app.function( volumes={VOL_MOUNT: volume}, gpu="A10G", timeout=7200, ) def pipeline(pretrain_steps=1000, sft_steps=500, gpu="a10g"): """Run pretraining then SFT in one go.""" print("=" * 60) print("PHASE 1: PRETRAINING") print("=" * 60) run_pretrain( steps=pretrain_steps, batch_size=8, grad_accum=4, seq_len=512, lr=3e-4, warmup=100, save_every=200, log_every=10, corpus="curated", resume=True, use_8bit_adam=True, ) print("\n" + "=" * 60) print("PHASE 2: SFT") print("=" * 60) run_sft( steps=sft_steps, batch_size=8, grad_accum=4, seq_len=768, lr=5e-5, warmup=20, save_every=100, log_every=10, use_8bit_adam=True, ) # ─── Local entry point for `modal run` ─────────────────────────────────────── @app.local_entrypoint() def main( command: str = "pipeline", steps: int = 1000, batch_size: int = 8, grad_accum: int = 4, seq_len: int = 512, lr: float = 3e-4, warmup: int = 100, save_every: int = 200, log_every: int = 10, corpus: str = "curated", resume: bool = True, use_8bit_adam: bool = True, gpu: str = "a10g", pretrain_steps: int = 1000, sft_steps: int = 500, ): if command == "seed": seed.remote() elif command == "pretrain": result = pretrain.remote( steps=steps, batch_size=batch_size, grad_accum=grad_accum, seq_len=seq_len, lr=lr, warmup=warmup, save_every=save_every, log_every=log_every, corpus=corpus, resume=resume, use_8bit_adam=use_8bit_adam, gpu=gpu, ) print(f"Pretraining final loss: {result}") elif command == "sft": result = sft.remote( steps=steps, batch_size=batch_size, grad_accum=grad_accum, seq_len=seq_len, lr=lr, warmup=warmup, save_every=save_every, log_every=log_every, use_8bit_adam=use_8bit_adam, gpu=gpu, ) print(f"SFT final loss: {result}") elif command == "download": download.remote() elif command == "pipeline": pipeline.remote(pretrain_steps=pretrain_steps, sft_steps=sft_steps, gpu=gpu) else: print(f"Unknown command: {command}") print("Available: upload, pretrain, sft, download, pipeline")