"""Forma del schedule WSD.""" from __future__ import annotations from enlace.config.schema import ScheduleConfig from enlace.train.schedules import build_lr_fn def _lr_fn(max_steps=1000, warmup=100, decay=200, min_ratio=0.0, base=1e-3): cfg = ScheduleConfig(kind="wsd", warmup_steps=warmup, decay_steps=decay, min_lr_ratio=min_ratio) return build_lr_fn(cfg, base, max_steps) def test_warmup_sube_linealmente_desde_arriba_de_cero(): lr = _lr_fn() assert lr(0) > 0.0 # el primer paso no se desperdicia con lr=0 assert lr(0) < lr(50) < lr(99) assert lr(99) == 1e-3 def test_la_fase_estable_es_plana(): lr = _lr_fn() assert lr(100) == lr(500) == lr(799) == 1e-3 def test_el_decay_baja_de_forma_monotona_hasta_cero(): lr = _lr_fn() valores = [lr(s) for s in range(800, 1000)] assert all(a >= b for a, b in zip(valores, valores[1:], strict=False)) assert valores[0] == 1e-3 assert valores[-1] < 1e-4 def test_min_lr_ratio_pone_un_piso(): lr = _lr_fn(min_ratio=0.1) assert lr(999) >= 1e-4 * 0.99 def test_sin_decay_el_lr_queda_plano_hasta_el_final(): lr = _lr_fn(decay=0) assert lr(999) == 1e-3 def test_extender_la_corrida_mueve_el_inicio_del_decay(): """Documenta la trampa: max_steps define dónde empieza a decaer el LR. Reanudar una corrida interrumpida exige el mismo max_steps; extenderla es una decisión distinta, que hay que tomar ramificando desde un checkpoint anterior al decay. """ corto = _lr_fn(max_steps=1000) largo = _lr_fn(max_steps=2000) assert corto(850) < corto(700) # ya está decayendo assert largo(850) == largo(700) # todavía en la fase estable