Byrne-100M-Ultra-MC / audit_architecture_paths.py
Quazim0t0's picture
Fix MC cache and attention masks; document generation checks and efficiency updates
c6835b3 verified
Raw History Blame Contribute Delete
2.17 kB
import json
from pathlib import Path
import torch
from config import SpikeWhaleConfig
from model_v2 import SpikeWhaleLM, reset_memory_cache
torch.set_num_threads(2)
torch.manual_seed(123)
c = torch.load('checkpoints/dpo_3200.pt', map_location='cpu', weights_only=False)
model = SpikeWhaleLM(SpikeWhaleConfig(**c['config'])).eval()
model.load_state_dict(c['model_state'], strict=True)
x = torch.tensor([[2, 41, 62, 81, 102, 121, 142, 161]])
changed = x.clone(); changed[:, 4:] += 31
def run(ids, mask=None):
reset_memory_cache(model)
return model(ids, attention_mask=mask).logits
results = {}
with torch.no_grad():
a, b = run(x), run(changed)
results['no_mask_future_prefix_delta'] = (a[:, :4]-b[:, :4]).abs().max().item()
mask = torch.ones_like(x)
a, b = run(x, mask), run(changed, mask)
results['binary_mask_future_prefix_delta'] = (a[:, :4]-b[:, :4]).abs().max().item()
try:
additive = torch.zeros(1,1,8,8).masked_fill(torch.triu(torch.ones(8,8,dtype=torch.bool),1), float('-inf'))
additive_logits = run(x, additive)
results['additive_binary_logits_delta'] = (additive_logits-a).abs().max().item()
results['additive_mask_error'] = None
except Exception as exc:
results['additive_mask_error'] = str(exc)
reset_memory_cache(model)
results['all_masked_loss'] = str(model(x, labels=torch.full_like(x,-100)).loss.item())
baseline = run(x)
model.train()
training = run(x)
results['train_eval_logits_delta'] = (baseline-training).abs().max().item()
from model_v2 import _masked_ce
empty_logits = torch.randn(5, 11, requires_grad=True)
empty_loss = _masked_ce(empty_logits, torch.full((5,), -100))
empty_loss.backward()
results['empty_ce_backward_finite'] = bool(torch.isfinite(empty_logits.grad).all())
assert results['no_mask_future_prefix_delta'] < 1e-4
assert results['binary_mask_future_prefix_delta'] < 1e-4
assert results['additive_mask_error'] is None
assert results['additive_binary_logits_delta'] < 1e-4
assert results['all_masked_loss'] == '0.0'
Path('architecture_path_audit.json').write_text(json.dumps(results, indent=2))
print(json.dumps(results, indent=2))