Download scripts/train.py from EndlessChasing/Mamb2_8B_Recall: direct link, hf CLI and curl.
- Browser
- Download file 12.2 kB
-
https://huggingface.co/EndlessChasing/Mamb2_8B_Recall/resolve/main/scripts/train.py
- Command line
-
hf download hf://EndlessChasing/Mamb2_8B_Recall/scripts/train.py
-
curl -L -o train.py https://huggingface.co/EndlessChasing/Mamb2_8B_Recall/resolve/main/scripts/train.py
12.2 kB
| #!/usr/bin/env python3 | |
| """Train the frozen-source Resurface control with the compressed arm's schedule.""" | |
| from __future__ import annotations | |
| import argparse | |
| import hashlib | |
| import json | |
| import math | |
| import os | |
| from pathlib import Path | |
| import sys | |
| import time | |
| import traceback | |
| import torch | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| from mamba2_recall import resurface_data as data, resurface_native as native | |
| from mamba2_recall import resurface_loss as objective | |
| from mamba2_recall import runtime | |
| STEPS = 1536 | |
| PROSE_MANIFEST_SHA = 'facb2ca461615a4199781bd21784d642d6674f5b862641b3b9edac3fb499b89d' | |
| PROSE_TOKENS_SHA = 'e54b02e5162e042a9cdd504f4eb1b1652724fb240bbc2c97608967aa26297233' | |
| def write_json(path, value): | |
| path = Path(path) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| temporary = path.with_suffix(path.suffix+'.tmp') | |
| temporary.write_text(json.dumps(value, indent=2, allow_nan=False)+'\n') | |
| temporary.replace(path) | |
| def lr_factor(j): | |
| return .1 + .9*(1+math.cos(math.pi*j/(STEPS-1)))/2 | |
| def schedule(): | |
| return torch.randperm(STEPS, generator=torch.Generator(device='cpu').manual_seed(2026092803)).tolist() | |
| def load_training_inputs(args): | |
| tokenizer = runtime.SentencePieceTokenizer(args.source_dir) | |
| manifest, _, examples = data.load_training(args.data_root, | |
| args.train_manifest_sha256, tokenizer) | |
| protocol_sha = data.sha_file(ROOT/'docs'/'PROTOCOL.md') | |
| if manifest['protocol_sha256'] != protocol_sha or len(examples) != STEPS: | |
| raise ValueError('Prepared numeric TRAIN split differs from frozen protocol') | |
| if data.sha_file(args.prose_manifest) != PROSE_MANIFEST_SHA: | |
| raise ValueError('Pinned WikiText TRAIN manifest differs') | |
| prose_manifest = json.loads(args.prose_manifest.read_text()) | |
| windows = torch.load(args.prose_tokens, map_location='cpu', weights_only=True) | |
| prose_file_sha = data.sha_file(args.prose_tokens) | |
| if (not prose_manifest.get('complete') or | |
| prose_manifest.get('training_tokens_file_sha256') != PROSE_TOKENS_SHA or | |
| tuple(windows.shape) != (448,2048) or windows.dtype != torch.int64 or | |
| sorted(prose_manifest['schedule']) != list(range(448)) or | |
| runtime.token_digest(windows.flatten().numpy()) != | |
| prose_manifest['training_tokens_sha256_int64le']): | |
| raise ValueError('Expected 448 disjoint, pinned prose TRAIN windows') | |
| binding = {'source_checkpoint_sha256': runtime.SOURCE_CHECKPOINT_SHA256, | |
| 'tokenizer_sha256': tokenizer.sha256, | |
| 'protocol_sha256': protocol_sha, | |
| 'train_manifest_sha256': args.train_manifest_sha256, | |
| 'prose_manifest_sha256': PROSE_MANIFEST_SHA, | |
| 'prose_tokens_sha256': prose_file_sha, | |
| 'prose_tokens_int64le_sha256': prose_manifest['training_tokens_sha256_int64le'], | |
| 'prose_file_exact_historical': prose_file_sha == PROSE_TOKENS_SHA, | |
| 'adapter': native.FORMAT, 'successful_updates': STEPS} | |
| return tokenizer, examples, windows, prose_manifest['schedule'], binding | |
| def pair_for(j, ordered, examples, windows, prose_order): | |
| row = examples[ordered[j]] | |
| full = torch.tensor(row['full_ids'], dtype=torch.long, device='cuda')[None] | |
| answer_mask = torch.zeros_like(full[:,:-1], dtype=torch.bool) | |
| answer_mask[:,row['ce_hidden_positions']] = True | |
| if int(answer_mask.sum()) != row['answer_target_count']: | |
| raise ValueError('Numeric answer suffix mask differs') | |
| prose_index = prose_order[j%448] | |
| prose_start = 512*((j//448)%4) | |
| prose = windows[prose_index,prose_start:prose_start+512].to('cuda')[None] | |
| if tuple(prose.shape) != (1,512): | |
| raise ValueError('Prose segment geometry differs') | |
| return {'id': row['id'], 'schedule_entry': ordered[j], | |
| 'prose_window': prose_index, 'prose_start': prose_start, | |
| 'mk_ids': full[:,:-1], 'mk_targets': full[:,1:], 'answer_mask': answer_mask, | |
| 'prose_ids': prose[:,:-1], 'prose_targets': prose[:,1:]} | |
| def optimizer_for(bank): | |
| mix = [p for name,p in bank.masters.items() if name.endswith(('.V_read','.g_read'))] | |
| router = [p for name,p in bank.masters.items() if name.endswith(('.router_w','.router_b'))] | |
| return torch.optim.AdamW([{'params':mix,'lr':1e-4,'base_lr':1e-4}, | |
| {'params':router,'lr':3e-4,'base_lr':3e-4}], | |
| betas=(.9,.999),eps=1e-8,weight_decay=0.) | |
| def attempt(bank, teacher, pair, optimizer, scaler, j): | |
| optimizer.zero_grad(set_to_none=True) | |
| for group in optimizer.param_groups: | |
| group['lr'] = group['base_lr']*lr_factor(j) | |
| hidden,gates = bank.forward_hidden(pair['mk_ids'], use_checkpoint=True) | |
| mk = objective.staged_loss_backward(hidden,pair['mk_targets'],bank.model.lm_head, | |
| gates=gates,scaler=scaler,answer_mask=pair['answer_mask'],ce_weight=1., | |
| kl_weight=0.,closure_weight=0.,chunk_tokens=64) | |
| del hidden,gates | |
| with torch.no_grad(): | |
| teacher_hidden = teacher.backbone(pair['prose_ids']) | |
| hidden,gates = bank.forward_hidden(pair['prose_ids'],use_checkpoint=True) | |
| prose = objective.staged_loss_backward(hidden,pair['prose_targets'],bank.model.lm_head, | |
| gates=gates,scaler=scaler,teacher_hidden=teacher_hidden, | |
| teacher_head=teacher.lm_head,ce_weight=.5,kl_weight=.5, | |
| closure_weight=3.,chunk_tokens=64) | |
| del hidden,gates,teacher_hidden | |
| scaler.unscale_(optimizer) | |
| gradients = [p.grad for p in bank.parameters()] | |
| if any(g is None for g in gradients): | |
| raise RuntimeError('Missing adapter gradient') | |
| overflow = (not all(bool(torch.isfinite(g).all()) for g in gradients) | |
| or not mk['scaled_hidden_gradient_finite'] | |
| or not prose['scaled_hidden_gradient_finite']) | |
| if overflow: | |
| scaler.update(new_scale=scaler.get_scale()*.5) | |
| magnitude = None | |
| else: | |
| magnitude = float(torch.nn.utils.clip_grad_norm_(bank.parameters(),1.,error_if_nonfinite=True)) | |
| scaler.step(optimizer) | |
| scaler.update() | |
| if any(not bool(torch.isfinite(p).all()) or not bool(torch.isfinite(p.half()).all()) | |
| for p in bank.parameters()): | |
| raise FloatingPointError('Nonfinite updated adapter master or FP16 cast') | |
| bank.assert_base_frozen() | |
| return {'mk_ce':mk['ce'], 'prose_ce':prose['ce'], | |
| 'prose_kl':prose['teacher_to_student_kl'], | |
| 'prose_closure':prose['closure'], 'overflow':overflow, | |
| 'gradient_norm_before_clip':magnitude, 'loss_scale':scaler.get_scale()} | |
| def save_checkpoint(path, bank, optimizer, scaler, binding, success, attempts): | |
| if path.exists(): | |
| raise FileExistsError(path) | |
| temporary = path.with_suffix('.pt.tmp') | |
| if temporary.exists(): | |
| raise FileExistsError(temporary) | |
| with temporary.open('wb') as stream: | |
| torch.save({'format':'MAMBA2_SOURCE_RESURFACE_CHECKPOINT_V1', | |
| 'binding':binding,'successful_updates':success,'attempts':attempts, | |
| 'masters':bank.state_dict(),'optimizer':optimizer.state_dict(), | |
| 'scaler':scaler.state_dict()}, stream) | |
| stream.flush(); os.fsync(stream.fileno()) | |
| temporary.replace(path) | |
| return {'path':str(path),'bytes':path.stat().st_size,'sha256':data.sha_file(path)} | |
| def main(): | |
| p = argparse.ArgumentParser(description=__doc__) | |
| p.add_argument('--source-dir',type=Path,required=True) | |
| p.add_argument('--data-root',type=Path,required=True) | |
| p.add_argument('--train-manifest-sha256',required=True) | |
| p.add_argument('--prose-manifest',type=Path,required=True) | |
| p.add_argument('--prose-tokens',type=Path,required=True) | |
| p.add_argument('--out-dir',type=Path,required=True) | |
| p.add_argument('--smoke',action='store_true',help='Discard one GPU update; no adapter export') | |
| args = p.parse_args() | |
| if args.out_dir.exists(): | |
| raise FileExistsError('New output directory required') | |
| args.out_dir.mkdir(parents=True) | |
| torch.set_num_threads(8) | |
| torch.manual_seed(2026092803); torch.cuda.manual_seed_all(2026092803) | |
| torch.backends.cuda.matmul.allow_tf32 = False | |
| torch.backends.cudnn.allow_tf32 = False | |
| torch.set_float32_matmul_precision('highest') | |
| report = {'format':'MAMBA2_SOURCE_RESURFACE_TRAIN_V1','complete':False, | |
| 'mode':'smoke' if args.smoke else 'formal', | |
| 'started_unix':time.time(),'source_sha256':data.sha_file(__file__), | |
| 'history':[],'checkpoints':[]} | |
| try: | |
| tokenizer, examples, windows, prose_order, binding = load_training_inputs(args) | |
| report['binding'] = binding | |
| write_json(args.out_dir/'report.json',report) | |
| student = runtime.load_source_model(args.source_dir) | |
| teacher = runtime.load_source_model(args.source_dir) | |
| if sum(p.numel() for p in student.parameters()) != 8236999680: | |
| raise RuntimeError('Wrong source model parameter count') | |
| bank = native.ResurfaceNative(student,'soft') | |
| if len(bank.masters)!=224 or sum(p.numel() for p in bank.parameters())!=1154104: | |
| raise RuntimeError('Wrong adapter geometry') | |
| optimizer = optimizer_for(bank) | |
| scaler = torch.amp.GradScaler('cuda',init_scale=1024.,growth_factor=2., | |
| backoff_factor=.5,growth_interval=2000) | |
| ordered = schedule() | |
| successful = attempts = 0 | |
| limit = 1 if args.smoke else STEPS | |
| while successful < limit: | |
| if attempts >= limit+8: | |
| raise RuntimeError('Overflow retry budget exceeded') | |
| pair = pair_for(successful,ordered,examples,windows,prose_order) | |
| start = time.monotonic() | |
| result = attempt(bank,teacher,pair,optimizer,scaler,successful) | |
| attempts += 1 | |
| if not result['overflow']: | |
| successful += 1 | |
| row = {'attempt':attempts,'successful_updates':successful, | |
| 'case_id':pair['id'],'schedule_entry':pair['schedule_entry'], | |
| 'prose_window':pair['prose_window'],'prose_start':pair['prose_start'], | |
| 'answer_targets':int(pair['answer_mask'].sum()), | |
| 'seconds':time.monotonic()-start,**result} | |
| report['history'].append(row) | |
| if not args.smoke and successful and successful%384==0 and not result['overflow']: | |
| report['checkpoints'].append(save_checkpoint(args.out_dir/f'checkpoint_{successful:04d}.pt', | |
| bank,optimizer,scaler,binding,successful,attempts)) | |
| if attempts%16==0 or successful==limit or result['overflow']: | |
| report.update(successful_updates=successful,attempts=attempts, | |
| gpu_memory=runtime.gpu_memory_receipt()) | |
| write_json(args.out_dir/'report.json',report) | |
| print(json.dumps({'updates':successful,'attempts':attempts, | |
| 'mk_ce':result['mk_ce'],'prose_ce':result['prose_ce'], | |
| 'overflow':result['overflow']}),flush=True) | |
| if not args.smoke: | |
| adapter_path = args.out_dir/'adapter_fp16.pt' | |
| exported = bank.export_fp16(adapter_path,binding=binding) | |
| if native.read_fp16(adapter_path,expected_binding=binding)['gate_mode']!='soft': | |
| raise RuntimeError('Export roundtrip differs') | |
| report['adapter'] = exported | |
| report['frozen_base_check'] = bank.assert_base_frozen() | |
| report['teacher_base_parameters_frozen'] = all( | |
| not p.requires_grad and p.grad is None for p in teacher.parameters()) | |
| report.update(complete=True,successful_updates=successful,attempts=attempts, | |
| gpu_memory=runtime.gpu_memory_receipt()) | |
| bank.close() | |
| except BaseException as error: | |
| report.update(error=repr(error),traceback=traceback.format_exc()) | |
| raise | |
| finally: | |
| report['finished_unix'] = time.time() | |
| write_json(args.out_dir/'report.json',report) | |
| if __name__=='__main__': | |
| main() | |