fix: resolve audited training, import, and serving bugs

- shard the Muon Newton-Schulz orthogonalization over the FSDP mesh instead of partial local slices
- import HF checkpoints faithfully: per-head RoPE permutation for q/k projections and qk-norm, qwen3, shared experts, and qk-norm before RoPE (changes numerics for existing use_qk_norm checkpoints)
- make preprocessing and resume self-contained: backfill realigned bucket keys by semantics (masks ones, rest zeros) and snapshot tokenizer files into every checkpoint
- keep RL consistent: sync the offline GRPO old_model each optimizer step and validate online strategies through a public one-off-rollout hook that leaves the replay cache untouched
- fix streaming serving: withhold partial tool-call prefixes with a stream-end flush, stream tool-call arguments from the raw source span, and terminate SSE frames with a blank line
- fix sampling semantics: capture logprobs before top-k/top-p mutate logits in place and detect greedy pipelines polymorphically instead of isinstance bookkeeping
This commit is contained in:
2026-09-03 20:27:41 +08:00
parent 7e98a419a7
commit 45cc048fe9
21 changed files with 834 additions and 72 deletions
+27
View File
@@ -45,6 +45,7 @@ class _RecordingRunner:
self._fresh = True
self.policy_version = result.policy_version
self.weight_updates = []
self.eval_calls = 0
def __call__(self, batch):
self.calls += 1
@@ -52,6 +53,12 @@ class _RecordingRunner:
self._fresh = False
return self.result, fresh
def evaluate(self, batch):
# Mirrors RolloutRunner.evaluate: one-off scoring that never
# touches the replay cache or freshness state.
self.eval_calls += 1
return self.result
def step(self):
self.step_calls += 1
@@ -360,6 +367,26 @@ def test_loss_is_differentiable_dpo(device):
assert has_grad
def test_validate_online_returns_none_without_runner(device):
strat = _make_grpo(device)
batch = {"input_ids": torch.randint(3, 200, (2, 4), device=device)}
assert strat.validate_online(batch) is None
def test_validate_online_uses_one_off_rollout_not_replay_cache(device):
strat = _make_grpo(device)
runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner)
out = strat.validate_online(
{"input_ids": torch.randint(3, 200, (2, 4), device=device)}
)
assert torch.isfinite(out["loss"]).item()
assert runner.eval_calls == 1
assert runner.calls == 0 # replay cache path untouched
def test_ref_model_not_updated_by_backward_dpo(device):
strat = _make_dpo(device)
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))