EndlessChasing's picture
Publish verified Mamba2-8B Resurface adapter and reproducibility package
5b7b27a verified
Raw History Blame Contribute Delete
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()