File size: 9,871 Bytes
5de8603 | 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 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 | #!/usr/bin/env python3
"""Teacher-forced held-out loss for continuous large-Talker checkpoints.
The fixed 50-row evaluation set is sourced from the benchmark's explicitly
held-out split. It is independent of the stage manifests in the production
plan. Each rank writes one Parquet shard and a small proof record.
"""
from __future__ import annotations
import argparse
import hashlib
import json
import math
import os
from pathlib import Path
import sys
import time
import numpy as np
SC = Path('/e/scratch/reformo/schuhmann1_moss')
HERE = Path(__file__).resolve().parent
V3 = Path('/e/home/jusers/schuhmann1/jupiter/m2_600m_tts_training/cascade_v3_timed')
M2 = SC / 'code/m2_20k_score'
SMALL = SC / 'code/small_tts'
SCORE = SC / 'out/m2_20k_score'
DEFAULT_PROMPTS = SCORE / 'eval/timed_prompts_v3.json'
DEFAULT_OUT = SC / 'out/m2_600m_ladder/caption_curriculum_v1/validation_loss'
def sha(path: Path) -> str:
digest = hashlib.sha256()
with path.open('rb') as stream:
for block in iter(lambda: stream.read(8 << 20), b''):
digest.update(block)
return digest.hexdigest()
def audit_heldout(prompts: list[dict], plan: dict) -> None:
assert len(prompts) == 50, f'expected fixed held-out set of 50 rows, got {len(prompts)}'
assert len({row['uid'] for row in prompts}) == len(prompts), 'duplicate held-out UID'
assert all(row.get('split') == 'heldout' for row in prompts), 'non-heldout prompt present'
# The ladder release contract defines tier 8 as a UID-disjoint held-out
# partition. Stage manifests contain only tier manifests, never heldout.
assert all('heldout' not in str(stage['manifest']).lower() for stage in plan['stages'])
assert all(row.get('codes') and row.get('ref_codes') for row in prompts), \
'held-out row lacks target or validated reference codes'
def load_checkpoint(model, checkpoint: Path, metadata: dict) -> str:
import torch
export = checkpoint / 'model_bf16.pt'
if export.exists():
assert metadata.get('export_sha256') in (None, sha(export)), 'export SHA mismatch'
state = torch.load(export, map_location='cpu', weights_only=True)
source = 'model_bf16.pt'
else:
state_path = checkpoint / 'training_state.pt'
assert state_path.stat().st_size == metadata['state_bytes']
assert sha(state_path) == metadata['state_sha256']
full = torch.load(state_path, map_location='cpu', weights_only=False)
state = {key[5:]: value for key, value in full['model'].items()
if key.startswith('base.')}
assert state and len(state) == len(full['model'])
source = 'training_state.pt:model'
del full
missing, unexpected = model.load_state_dict(state, strict=True)
assert not missing and not unexpected
model.tie_weights()
del state
return source
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument('--checkpoint', required=True)
parser.add_argument('--prompts', default=str(DEFAULT_PROMPTS))
parser.add_argument('--output-root', default=str(DEFAULT_OUT))
parser.add_argument('--limit', type=int, help='deterministic prefix for a canary')
args = parser.parse_args()
checkpoint = Path(args.checkpoint).resolve()
metadata_path = checkpoint / 'COMPLETE.json'
metadata = json.loads(metadata_path.read_text())
assert int(metadata['step']) > 0 and metadata['world'] == 128
assert (checkpoint / 'training_state.pt').stat().st_size == metadata['state_bytes']
plan_path = HERE / 'plan_caption_production_32nodes.json'
plan = json.loads(plan_path.read_text())
prompts_path = Path(args.prompts).resolve()
prompts = json.loads(prompts_path.read_text())
audit_heldout(prompts, plan)
if args.limit is not None:
assert 1 <= args.limit <= len(prompts)
prompts = prompts[:args.limit]
rank = int(os.environ.get('SLURM_PROCID', '0'))
world = int(os.environ.get('SLURM_NTASKS', '1'))
local = int(os.environ.get('SLURM_LOCALID', '0'))
assert world in (1, 4), f'evaluation expects one or four ranks, got {world}'
for candidate in (str(HERE), str(V3), str(M2), str(SMALL)):
while candidate in sys.path:
sys.path.remove(candidate)
sys.path[:0] = [str(HERE), str(V3), str(M2), str(SMALL)]
import torch
from transformers import AutoTokenizer
from packing import ScorePacker
import moss_small
import va_loss
from large_talker import build_fresh
device_index = 0 if torch.cuda.device_count() == 1 else local
torch.cuda.set_device(device_index)
device = torch.device('cuda', device_index)
torch.set_num_threads(4)
schema_path = SCORE / 'score_schema.json'
schema = json.loads(schema_path.read_text())
model, config = build_fresh(schema, log=lambda _: None)
checkpoint_source = load_checkpoint(model, checkpoint, metadata)
model = model.to(device, dtype=torch.bfloat16).eval()
_, _, Processor = moss_small.export_classes()
processor = Processor(tokenizer=AutoTokenizer.from_pretrained(
moss_small.SFT3, trust_remote_code=True, local_files_only=True),
audio_tokenizer=None, model_config=config)
packer = ScorePacker(processor, config, schema)
tasks = [(index, mode) for index in range(len(prompts))
for mode in ('instruction', 'reference')]
mine = [task for position, task in enumerate(tasks) if position % world == rank]
started = time.monotonic()
rows = []
per_channel_numerator = np.zeros(13, dtype=np.float64)
per_channel_denominator = np.zeros(13, dtype=np.float64)
for index, mode in mine:
prompt = prompts[index]
reference = prompt['ref_codes'] if mode == 'reference' else None
example = packer.pack_mode(prompt, prompt['codes'], mode, reference, generation=False)
batch = packer.collate([example])
inputs = {
'input_ids': batch['input_ids'].to(device),
'attention_mask': batch['attention_mask'].to(device),
'labels': batch['labels'].to(device),
'score_conditioning': tuple(value.to(device) for value in batch['score_conditioning']),
}
with torch.inference_mode(), torch.autocast('cuda', dtype=torch.bfloat16):
hidden = model(input_ids=inputs['input_ids'], attention_mask=inputs['attention_mask'],
score_conditioning=inputs['score_conditioning'],
use_cache=False).last_hidden_state
loss, per = va_loss.compute_supervised_loss_from_hidden(
model, global_hidden_states=hidden, labels=inputs['labels'],
channelwise_loss_weight=[1.0] + [32.0 / 12.0] * 12,
return_per_channel=True)
frame_count = int(batch['labels'][:, :, 1].ge(0).sum().item())
sample_count = 1
assert frame_count == len(prompt['codes'])
per_values = per.detach().float().cpu().numpy().reshape(-1)
assert len(per_values) == 13 and np.isfinite(per_values).all()
counts = np.asarray([frame_count + sample_count] + [frame_count] * 12,
dtype=np.float64)
per_channel_numerator += per_values * counts
per_channel_denominator += counts
rows.append({
'uid': prompt['uid'], 'split': prompt['split'], 'source': prompt['source'],
'lang': prompt['lang'], 'mode': mode, 'checkpoint_step': int(metadata['step']),
'frames': frame_count, 'samples': sample_count,
'loss': float(loss.detach().float().cpu()),
'per_channel_loss_json': json.dumps(per_values.tolist()),
})
del hidden, batch, example
if len(rows) % 10 == 0:
print('VALIDATION_ROW', rank, len(rows), flush=True)
import pyarrow as pa
import pyarrow.parquet as pq
destination = Path(args.output_root) / checkpoint.name / ('canary' if args.limit else 'full')
destination.mkdir(parents=True, exist_ok=True)
parquet_path = destination / f'loss.rank{rank:03d}.parquet'
temporary = parquet_path.with_suffix('.parquet.tmp')
pq.write_table(pa.Table.from_pylist(rows), temporary, compression='zstd')
temporary.replace(parquet_path)
channels = np.divide(per_channel_numerator, per_channel_denominator,
out=np.zeros_like(per_channel_numerator),
where=per_channel_denominator > 0)
weights = np.asarray([1.0] + [32.0 / 12.0] * 12)
proof = {
'status': 'PASS', 'rank': rank, 'world': world,
'checkpoint': str(checkpoint), 'checkpoint_step': int(metadata['step']),
'checkpoint_contract': metadata['contract'],
'checkpoint_source': checkpoint_source,
'heldout': True, 'heldout_rows_in_input': len(prompts), 'task_rows': len(rows),
'prompts_path': str(prompts_path), 'prompts_sha256': sha(prompts_path),
'schema_sha256': sha(schema_path), 'production_plan_sha256': sha(plan_path),
'evaluator_sha256': sha(Path(__file__)),
'per_channel_loss': channels.tolist(),
'per_channel_numerator': per_channel_numerator.tolist(),
'per_channel_denominator': per_channel_denominator.tolist(),
'validation_loss': float(np.dot(channels, weights) / weights.sum()),
'parquet': str(parquet_path), 'parquet_sha256': sha(parquet_path),
'elapsed_seconds': time.monotonic() - started,
'slurm_job_id': os.environ.get('SLURM_JOB_ID'),
}
proof_path = destination / f'rank{rank:03d}.PASS.json'
temporary_proof = proof_path.with_suffix('.json.tmp')
temporary_proof.write_text(json.dumps(proof, indent=2) + '\n')
temporary_proof.replace(proof_path)
print('VALIDATION_RANK_PASS', rank, len(rows), proof['validation_loss'], flush=True)
if __name__ == '__main__':
main()
|