"""Same global sample stream across ranks, without dropped tail targets.""" import hashlib import json import numpy as np def benchmark_phases(world): assert world in (8, 16) return [{'name': 'strong', 'offset': 0, 'global_batch': 256, 'updates': 40, 'warmup': 10}, {'name': 'weak', 'offset': 10240, 'global_batch': world * 32, 'updates': 12, 'warmup': 4, 'replicated_base_world': 8}] def step_indices(phase, step, rank, world, samples): if phase.get('replicated_base_world'): # Benchmark only: duplicate the same eight real rank streams on the # second pair of nodes, giving exactly equal per-GPU input/token load. base = phase['replicated_base_world'] local_count = phase['global_batch'] // world start = phase['offset'] + step * base * local_count return list(range(start + rank % base, start + base * local_count, base)) start = phase['offset'] + step * phase['global_batch'] end = min(samples, start + phase['global_batch']) if start >= samples: return [] return list(range(start + rank, end, world)) def microbatches(examples, max_padded_tokens=12000, max_examples=16): """Length sort only within one already fixed optimizer batch.""" examples = sorted(examples, key=lambda ex: len(ex['input_ids'])) batches, current = [], [] for row in examples: n = len(row['input_ids']) if current and (n * (len(current) + 1) > max_padded_tokens or len(current) >= max_examples): batches.append(current); current = [] current.append(row) if current: batches.append(current) return batches def dynamic_weights(local_frames, local_samples, global_frames, global_samples, world): """Preserve one global per-channel token mean across microbatches/ranks. The inherited loss normalizes by sum(weights), so the caller multiplies its return value by returned scale. DDP subsequently averages gradients. This is algebraically the ordinary inherited objective on the global batch. """ weights = [world * (local_frames + local_samples) / (global_frames + global_samples)] weights += [(32. / 12) * world * local_frames / global_frames] * 12 return weights, sum(weights) / 33. def check_plan(samples=16384): unions = {} for world in (8, 16): phases = benchmark_phases(world) for phase in phases: seen = [] for step in range(phase['updates']): global_ids = sorted(i for rank in range(world) for i in step_indices(phase, step, rank, world, samples)) if phase['name'] == 'weak': start = phase['offset'] + step * 256 expect = sorted(list(range(start, start + 256)) * (world // 8)) else: expect = list(range(phase['offset'] + step * phase['global_batch'], phase['offset'] + (step + 1) * phase['global_batch'])) assert global_ids == expect seen.extend(global_ids) assert len(seen) == len(set(seen)) * (world // 8 if phase['name'] == 'weak' else 1) unions[(world, phase['name'])] = seen assert unions[(8, 'strong')] == unions[(16, 'strong')] for step in range(12): for rank in range(16): assert step_indices(benchmark_phases(16)[1], step, rank, 16, samples) == step_indices(benchmark_phases(8)[1], step, rank % 8, 8, samples) # Test mean/gradient coefficients independently, including unequal lengths. rng = np.random.default_rng(71) channels = rng.normal(size=(263, 13)) frame_groups = [13, 49, 81, 120] sample_groups = [1, 3, 5, 7] expected = channels.mean(axis=0) actual = np.zeros(13) start = 0 for frames, sample_count in zip(frame_groups, sample_groups): weights, scale = dynamic_weights(frames, sample_count, sum(frame_groups), sum(sample_groups), 4) # Audio-channel linear coefficients recover the exact global mean after DDP. audio_coefficient = weights[1] / sum(weights) * scale / 4 actual[1:] += channels[start:start + frames, 1:].mean(axis=0) * audio_coefficient start += frames assert np.allclose(actual[1:], expected[1:] * (32 / 12) / 33) return {'status': 'PASS', 'strong_unique_samples': len(unions[(8, 'strong')]), 'strong_union_sha256': hashlib.sha256(np.asarray(unions[(8, 'strong')], dtype='