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
+43
View File
@@ -3,6 +3,7 @@
import torch
from astrai.inference.runtime.sample import (
BaseSamplingStrategy,
FrequencyPenaltyStrategy,
SamplingPipeline,
TemperatureStrategy,
@@ -295,3 +296,45 @@ def test_greedy_respects_frequency_penalty():
)
# Token 0 saw four occurrences: 5 - 2*4 < 4, so the argmax flips.
assert penalized.tolist() == [1]
class _ArgmaxMovingStrategy(BaseSamplingStrategy):
"""Custom strategy that can move the argmax — must disable greedy."""
def apply(
self, logits, filter_value=-float("inf"), input_ids=None, input_mask=None
):
return torch.roll(logits, shifts=1, dims=-1)
def test_greedy_detection_is_polymorphic():
"""Greedy detection asks strategies polymorphically, no isinstance."""
base = [TemperatureStrategy(0.0), TopKStrategy(50), TopPStrategy(0.9)]
assert SamplingPipeline(list(base)).is_greedy is True
assert SamplingPipeline(base + [FrequencyPenaltyStrategy(0.5)]).is_greedy is False
# A custom argmax-moving strategy disables greedy even though the
# pipeline contains a greedy temperature — this is what isinstance
# bookkeeping in the old implementation could not see.
assert SamplingPipeline(base + [_ArgmaxMovingStrategy()]).is_greedy is False
def test_greedy_detection_position_independent():
"""Greedy temperature anywhere in the pipeline is detected."""
pipeline = SamplingPipeline([TopKStrategy(50), TemperatureStrategy(0.0)])
assert pipeline.is_greedy is True
def test_greedy_detection_composes_across_nested_pipelines():
"""A nested pipeline participates through the same interface."""
inner = SamplingPipeline([TemperatureStrategy(0.0), TopKStrategy(20)])
assert inner.is_greedy is True
assert SamplingPipeline([TopPStrategy(0.9), inner]).is_greedy is True
assert SamplingPipeline([inner, FrequencyPenaltyStrategy(0.5)]).is_greedy is False
def test_nongreedy_temperature_is_not_greedy():
pipeline = SamplingPipeline(
[TemperatureStrategy(0.7), TopKStrategy(0), TopPStrategy(1.0)]
)
assert pipeline.is_greedy is False
+37
View File
@@ -565,3 +565,40 @@ def test_parser_uses_token_ids_for_detection():
parser = TokenIdParser()
parser.feed("hello", current_token_ids=[1, 999, 3])
assert parser.has_tool_calls
def test_streaming_partial_name_prefix_never_leaks_into_content():
parser = SimpleJsonToolParser()
parts = ["Hello ", '{"', '{"n', '{"na', '{"name"']
emitted = []
body = ""
for part in parts:
body += part
for d in parser.feed(body):
if "content" in d:
emitted.append(d["content"])
assert "".join(emitted) == "Hello "
def test_finalize_flushes_withheld_plain_json_content():
parser = SimpleJsonToolParser()
text = 'Answer: {"price": 1}'
deltas = parser.feed(text)
streamed = "".join(d["content"] for d in deltas if "content" in d)
flushed = parser.finalize(text)
joined = streamed + "".join(d["content"] for d in flushed if "content" in d)
assert joined == text
assert not parser.has_tool_calls
assert parser.finalize(text) == []
def test_streaming_args_concat_matches_parse_complete():
parser = SimpleJsonToolParser()
# Compact spacing: json.dumps would re-space this and desync the
# streamed arguments diff.
text = '{"name": "get_weather","arguments": {"city":"Beijing","unit":"c"}}'
_, args_chunks = _simulate_streaming(parser, text)
streamed = "".join(args_chunks)
completed = parser.parse_complete(text)["tool_calls"][0]["function"]["arguments"]
assert streamed == completed
assert streamed == '"city":"Beijing","unit":"c"'