File size: 6,606 Bytes
cd9b2d8 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 | #!/usr/bin/env python3
"""Four independent inference ranks; immutable resumable WebDataset caches."""
import argparse
import io
import json
import os
import tarfile
import time
import hashlib
from pathlib import Path
import numpy as np
import torch
from study_paths import ROOT, read, write
from backbones import FrozenEncoder
from study_data import SourceData
def member(tf, name, payload):
offset = tf.offset + 512
info = tarfile.TarInfo(name)
info.size, info.mtime = len(payload), 0
tf.addfile(info, io.BytesIO(payload))
return offset, len(payload)
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--model', required=True)
ap.add_argument('--smoke', action='store_true')
ap.add_argument('--domains', nargs='+')
args = ap.parse_args()
cfg = read(ROOT / 'study.json')
spec = next(m for m in cfg['models'] if m['id'] == args.model)
rank, world = int(os.environ.get('LOCAL_RANK', 0)), int(os.environ.get('WORLD_SIZE', 1))
torch.set_num_threads(2)
torch.cuda.set_device(rank)
encoder = FrozenEncoder(spec, torch.device('cuda', rank))
out = ROOT / ('smoke' if args.smoke else 'features') / args.model
out.mkdir(parents=True, exist_ok=True)
domains = args.domains or ['ladder', 'p3_train', 'p3_validation', 'p3_test', 'gemini', *cfg['benchmarks']]
for domain in domains:
data = SourceData(domain)
positions = list(range(rank, len(data), world))
if args.smoke:
positions = positions[:4]
completed = 0
started = time.monotonic()
for chunk_number, offset in enumerate(range(0, len(positions), 512)):
name = f'{domain}-rank{rank}-{chunk_number:05d}'
destination = out / (name + '.tar')
manifest = out / (name + '.jsonl')
commit = out / (name + '.commit.json')
chunk = positions[offset:offset + 512]
if commit.exists():
completed += len(chunk)
continue
rows = []
temp = destination.with_suffix('.tar.incomplete')
with tarfile.open(temp, 'w', format=tarfile.USTAR_FORMAT) as tf:
for step in range(0, len(chunk), spec['batch']):
items = [data[p] for p in chunk[step:step + spec['batch']]]
waves = [row[1] for row in items]
encoded = encoder.encode(waves)
for local_index, ((key, wave, meta, targets), features) in enumerate(zip(items, encoded)):
signature = hashlib.sha256(key.encode()).hexdigest()
buffer = io.BytesIO()
np.savez(buffer, **features, **targets)
payload = buffer.getvalue()
position, size = member(tf, signature + '.npz', payload)
rows.append({**meta, 'feature_key': key, 'source_position': int(chunk[step + local_index]),
'feature_tar': str(destination), 'feature_offset': position, 'feature_size': size,
'feature_sha256': hashlib.sha256(payload).hexdigest(),
'decoded_waveform_sha256': hashlib.sha256(wave.astype('<f4').tobytes()).hexdigest(),
'native_dim': spec['native_dim'], 'duration_s': float(features['duration_s']),
'frame_count': len(features['frame_features'])})
temp.replace(destination)
temporary = manifest.with_suffix('.jsonl.incomplete')
temporary.write_text(''.join(json.dumps(r, ensure_ascii=False) + '\n' for r in rows))
temporary.replace(manifest)
write(commit, {'clips': len(rows), 'tar': str(destination), 'index': str(manifest),
'frozen_backbone': spec, 'feature_schema': 'pooled-native+temporal-fixed64-v1'})
completed += len(rows)
write(out / f'{domain}-rank{rank}-progress.json', {'completed': completed, 'total': len(positions),
'clips_per_second': completed / max(.001, time.monotonic() - started), 'model': args.model, 'domain': domain})
print(args.model, domain, rank, completed, '/', len(positions), flush=True)
write(out / f'{domain}-rank{rank}-COMPLETE.json', {'clips': completed, 'total': len(positions), 'world': world})
if args.smoke:
# Exercise real cached Ladder/P3 targets, the loss, and a head backward
# pass before committing a full encoder-cache allocation.
import sys
from study_paths import CODE, RELEASE
from cache_dataset import collate
from probe_model import Probe
sys.path.insert(0, str(CODE))
from train_layered_curriculum import loss_terms
samples = []
for domain in ('ladder', 'p3_train'):
source = SourceData(domain)
_, wave, _, targets = source[rank]
features = encoder.encode([wave])[0]
ticks = (np.arange(len(targets['frame'])) + .5) / 50
native = features['frame_features']
samples.append({**targets, 'embedding': np.pad(features['embedding'].astype(np.float32)[:256],
(0, max(0, 256-len(features['embedding'])))),
'frame_features': np.stack([np.interp(ticks, features['frame_times_s'], native[:, i])
for i in range(64)], axis=1).astype(np.float32),
'feature_key': domain + '-contract'})
batch = {k: v.to(encoder.device) if torch.is_tensor(v) else v for k,v in collate(samples).items()}
classes = read(RELEASE / 'classes.json')['names']
probe = Probe(len(classes)).to(encoder.device)
result = probe(batch['embedding'], batch['frame_features'], batch['event_starts'], batch['event_ends'])
loss = sum(loss_terms(result, batch, torch.ones(len(classes), device=encoder.device)).values())
loss.backward()
if not torch.isfinite(loss) or not any(p.grad is not None and torch.isfinite(p.grad).all() for p in probe.parameters()):
raise RuntimeError('Probe backward contract failed')
assert not any(p.requires_grad or p.grad is not None for p in encoder.model.parameters())
write(out / f'CONTRACT-rank{rank}.json', {'loss': float(loss.detach()), 'finite_gradient': True,
'frozen_backbone': True, 'classes': len(classes)})
if __name__ == '__main__':
main()
|