#!/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('