"""Runtime checks before training: FlashAttention kernel vs SDPA agreement, and whether a 65,536-token microbatch fits without activation checkpointing (sets the engine's checkpointing threshold).""" import gc, time import numpy as np import torch from transformers import AutoModelForSequenceClassification, AutoTokenizer from common import * def step(model, ids, checkpoint): if checkpoint: model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={'use_reentrant': False}) else: model.gradient_checkpointing_disable() torch.cuda.reset_peak_memory_stats(); torch.cuda.synchronize(); t = time.time() with torch.autocast('cuda', dtype=torch.bfloat16): logits = model(input_ids=ids, attention_mask=torch.ones_like(ids)).logits.float() loss = torch.nn.functional.cross_entropy(logits, torch.zeros(ids.shape[0], dtype=torch.long, device='cuda')) loss.backward(); torch.cuda.synchronize() out = {'shape': list(ids.shape), 'checkpoint': checkpoint, 'seconds': time.time()-t, 'peak_gb': torch.cuda.max_memory_allocated()/1e9, 'finite': bool(torch.isfinite(logits).all())} model.zero_grad(set_to_none=True) return out def main(): if (DATA/'smoke.json').exists(): print('smoke done'); return path = init_path(); tok = AutoTokenizer.from_pretrained(path) sample = '### PROPOSED TOOL CALL\ntool: Bash\nargs: rm -rf node_modules dist && npm install\n\n### USER REQUEST\nClean up the build artifacts and reinstall dependencies.\n\n### AGENT HISTORY\n[1] Bash(ls)\n-> node_modules dist package.json' x = tok([sample, sample + '\n' + 'filler ' * 700], return_tensors='pt', padding=True).to('cuda') fa = AutoModelForSequenceClassification.from_pretrained(path, dtype=torch.bfloat16, attn_implementation=ATTENTION, allow_all_kernels=True).cuda().eval() sd = AutoModelForSequenceClassification.from_pretrained(path, dtype=torch.bfloat16, attn_implementation='sdpa').cuda().eval() with torch.inference_mode(): a = fa(**x).logits.float(); b = sd(**x).logits.float() delta = float((a-b).abs().max()) assert torch.isfinite(a).all() and delta < 0.1, f'flash vs sdpa logits differ by {delta}' del fa, sd; gc.collect(); torch.cuda.empty_cache() model = AutoModelForSequenceClassification.from_pretrained(path, dtype=torch.float32, attn_implementation=ATTENTION, allow_all_kernels=True).cuda().train() g = torch.Generator().manual_seed(0) runs = [] for shape, ckpt in [((128, 512), False), ((1, MAX_LEN), False), ((1, MAX_LEN), True), ((4, 16384), False)]: ids = torch.randint(1000, 50000, shape, generator=g).cuda(); ids[:, 0] = 50281; ids[:, -1] = 50282 try: runs.append(step(model, ids, ckpt)) except torch.OutOfMemoryError as e: runs.append({'shape': list(shape), 'checkpoint': ckpt, 'oom': True}); model.zero_grad(set_to_none=True) gc.collect(); torch.cuda.empty_cache() print(runs[-1], flush=True) full = next(r for r in runs if r['shape'] == [1, MAX_LEN] and not r['checkpoint']) # Leave >= 20 GB headroom for optimizer state, the data loader and fragmentation. ckpt_tokens = MAX_LEN + 2048 if (not full.get('oom') and full['peak_gb'] < 70) else 16384 # +2048: collate pads rows to a multiple of 8 result = {'flash_vs_sdpa_max_logit_delta': delta, 'runs': runs, 'ckpt_tokens': ckpt_tokens, 'device': torch.cuda.get_device_name(), 'torch': torch.__version__} atomic_json(DATA/'smoke.json', result); event('smoke', **result) if __name__ == '__main__': main()