Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
ChristophSchuhmann's picture
Document architecture, prompts, code, and full run statistics
d911efa verified
Raw History Blame Contribute Delete
5.04 kB
"""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='<i8').tobytes()).hexdigest(),
'worlds': [8, 16], 'global_batch_strong': 256, 'per_gpu_weak': 32,
'strong_warmup_updates': 10, 'strong_steady_updates': 30, 'weak_warmup_updates': 4, 'weak_steady_updates': 8,
'weak_input_rule': 'Exact same eight real per-GPU streams, replicated once across the additional eight GPUs at four nodes. Benchmark-only repetitions, explicitly counted as presentations.',
'loss_weighting': 'Exact global per-channel supervised-token mean, independent of rank and microbatch partition.'}
if __name__ == '__main__': print(json.dumps(check_plan(), indent=2))