diff --git a/python/sglang/test/test_deterministic.py b/python/sglang/test/test_deterministic.py index 20d45c3ef..17f5495ab 100644 --- a/python/sglang/test/test_deterministic.py +++ b/python/sglang/test/test_deterministic.py @@ -346,29 +346,61 @@ class TokenIdsAndLogprobs: print(f"✅ Logprobs match:", a.logprobs[:5]) else: print(f"❌ Logprobs mismatch") - # Only print last 10 elements for readability - n_show = 10 - a_show = a.logprobs[-n_show:] - b_show = b.logprobs[-n_show:] - print( - " A: ... ", - [f"{x:.10f}" if x is not None else "None" for x in a_show], - f"({len(a.logprobs)} total)" if len(a.logprobs) > n_show else "", - ) - print( - " B: ... ", - [f"{x:.10f}" if x is not None else "None" for x in b_show], - f"({len(b.logprobs)} total)" if len(b.logprobs) > n_show else "", - ) - diff = [ - abs(x - y) if x is not None else float("nan") - for x, y in zip(a.logprobs, b.logprobs) - ] - print( - " Diff:", - [f"{x:.10e}" for x in diff[-n_show:]], - f"... ({len(diff)} total)" if len(diff) > n_show else "", - ) + + # Find first divergent position + first_div = None + for idx, (la, lb) in enumerate(zip(a.logprobs, b.logprobs)): + if la != lb: + first_div = idx + break + + n_show = 5 + if first_div is not None: + print(f" First divergence at position {first_div}/{len(a.logprobs)}") + # Show n_show elements starting from the divergent point + a_show = a.logprobs[first_div : first_div + n_show] + b_show = b.logprobs[first_div : first_div + n_show] + diff_show = [ + abs(x - y) if x is not None and y is not None else float("nan") + for x, y in zip(a_show, b_show) + ] + pos_range = f"[{first_div}:{first_div + len(a_show)}]" + label_width = len(f"A {pos_range}") + print( + f" A {pos_range}: ", + [f"{x:.10f}" if x is not None else "None" for x in a_show], + ) + print( + f" B {pos_range}: ", + [f"{x:.10f}" if x is not None else "None" for x in b_show], + ) + print( + f" {'Diff':<{label_width}}: ", + [f"{x:.10e}" for x in diff_show], + ) + else: + # Fallback to tail (shouldn't happen if logprobs_match is False) + a_show = a.logprobs[-n_show:] + b_show = b.logprobs[-n_show:] + diff_show = [ + abs(x - y) if x is not None and y is not None else float("nan") + for x, y in zip(a_show, b_show) + ] + print( + " A: ... ", + [f"{x:.10f}" if x is not None else "None" for x in a_show], + f"({len(a.logprobs)} total)" if len(a.logprobs) > n_show else "", + ) + print( + " B: ... ", + [f"{x:.10f}" if x is not None else "None" for x in b_show], + f"({len(b.logprobs)} total)" if len(b.logprobs) > n_show else "", + ) + print( + " Diff: ... ", + [f"{x:.10e}" for x in diff_show], + f"({len(a.logprobs)} total)" if len(a.logprobs) > n_show else "", + ) # Compute KL-divergence using K3 approximation # KL(P||Q) ≈ (exp(log(P) - log(Q)) - 1) - (log(P) - log(Q))