[Auto Sync] Update test_deterministic.py (20260124) (#17665)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com> Co-authored-by: Jiayi Yuan <34369239+jy-yuan@users.noreply.github.com>
This commit is contained in:
co-authored by
github-actions[bot]
Jiayi Yuan
parent
a6280b2a23
commit
0834f9afeb
@@ -343,19 +343,19 @@ class TokenIdsAndLogprobs:
|
|||||||
print(f"✅ Logprobs match:", a.logprobs[:5])
|
print(f"✅ Logprobs match:", a.logprobs[:5])
|
||||||
else:
|
else:
|
||||||
print(f"❌ Logprobs mismatch")
|
print(f"❌ Logprobs mismatch")
|
||||||
# Only print first 5 elements for readability
|
# Only print last 10 elements for readability
|
||||||
n_show = 5
|
n_show = 10
|
||||||
a_show = a.logprobs[:n_show]
|
a_show = a.logprobs[-n_show:]
|
||||||
b_show = b.logprobs[:n_show]
|
b_show = b.logprobs[-n_show:]
|
||||||
print(
|
print(
|
||||||
" A: ",
|
" A: ... ",
|
||||||
[f"{x:.10f}" if x is not None else "None" for x in a_show],
|
[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 "",
|
f"({len(a.logprobs)} total)" if len(a.logprobs) > n_show else "",
|
||||||
)
|
)
|
||||||
print(
|
print(
|
||||||
" B: ",
|
" B: ... ",
|
||||||
[f"{x:.10f}" if x is not None else "None" for x in b_show],
|
[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 "",
|
f"({len(b.logprobs)} total)" if len(b.logprobs) > n_show else "",
|
||||||
)
|
)
|
||||||
diff = [
|
diff = [
|
||||||
abs(x - y) if x is not None else float("nan")
|
abs(x - y) if x is not None else float("nan")
|
||||||
@@ -363,7 +363,7 @@ class TokenIdsAndLogprobs:
|
|||||||
]
|
]
|
||||||
print(
|
print(
|
||||||
" Diff:",
|
" Diff:",
|
||||||
[f"{x:.10e}" for x in diff[:n_show]],
|
[f"{x:.10e}" for x in diff[-n_show:]],
|
||||||
f"... ({len(diff)} total)" if len(diff) > n_show else "",
|
f"... ({len(diff)} total)" if len(diff) > n_show else "",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -528,23 +528,25 @@ def test_deterministic(args):
|
|||||||
# Flush cache first to make sure there is no cache hit from previous tests
|
# Flush cache first to make sure there is no cache hit from previous tests
|
||||||
flush_response = requests.post(f"http://{args.host}:{args.port}/flush_cache")
|
flush_response = requests.post(f"http://{args.host}:{args.port}/flush_cache")
|
||||||
|
|
||||||
print(f"Step 1: Generating random 64 token IDs...")
|
prefix_len = 100
|
||||||
|
print(f"Step 1: Generating random {prefix_len} token IDs...")
|
||||||
# Use a reasonable token ID range (e.g., 1-50000 for most tokenizers)
|
# Use a reasonable token ID range (e.g., 1-50000 for most tokenizers)
|
||||||
# Avoid special tokens like 0 (padding), 1 (BOS), 2 (EOS)
|
# Avoid special tokens like 0 (padding), 1 (BOS), 2 (EOS)
|
||||||
# set seed for random.randint
|
# set seed for random.randint
|
||||||
random.seed(42)
|
random.seed(42)
|
||||||
initial_token_ids = [random.randint(100, 50000) for _ in range(64)]
|
initial_token_ids = [random.randint(100, 50000) for _ in range(prefix_len)]
|
||||||
|
|
||||||
print(f"✓ Using {len(initial_token_ids)} initial tokens")
|
print(f"✓ Using {len(initial_token_ids)} initial tokens")
|
||||||
print(f" Initial token IDs: {initial_token_ids}")
|
print(f" Initial token IDs: {initial_token_ids}")
|
||||||
|
|
||||||
|
num_tokens_to_generate = 2
|
||||||
print(
|
print(
|
||||||
f"\nStep 2: Generating 2 tokens from {len(initial_token_ids)} token prefix..."
|
f"\nStep 2: Generating {num_tokens_to_generate} tokens from {len(initial_token_ids)} token prefix..."
|
||||||
)
|
)
|
||||||
first_response = send_single(
|
first_response = send_single(
|
||||||
args,
|
args,
|
||||||
input_ids=initial_token_ids,
|
input_ids=initial_token_ids,
|
||||||
max_new_tokens=100,
|
max_new_tokens=num_tokens_to_generate,
|
||||||
return_full_response=True,
|
return_full_response=True,
|
||||||
)
|
)
|
||||||
first_output_text = first_response["text"]
|
first_output_text = first_response["text"]
|
||||||
@@ -558,11 +560,11 @@ def test_deterministic(args):
|
|||||||
print(f' Output text: "{first_output_text}"')
|
print(f' Output text: "{first_output_text}"')
|
||||||
|
|
||||||
print(
|
print(
|
||||||
f"\nStep 3: Generating with radix cache (164 tokens prefill, should hit > 128 tokens cache, based on page size)..."
|
f"\nStep 3: Generating with radix cache ({len(initial_token_ids + first_output_token_ids[:-1])} tokens prefill, should hit cache based on page size)..."
|
||||||
)
|
)
|
||||||
prefix_token_ids = initial_token_ids + first_output_token_ids[:-1]
|
prefix_token_ids = initial_token_ids + first_output_token_ids[:-1]
|
||||||
print(
|
print(
|
||||||
f" Prefix: {len(initial_token_ids)} initial + 64 generated = {len(prefix_token_ids)} tokens"
|
f" Prefix: {len(initial_token_ids)} initial + 1 generated = {len(prefix_token_ids)} tokens"
|
||||||
)
|
)
|
||||||
print(f"Using Prompt: {prefix_token_ids}")
|
print(f"Using Prompt: {prefix_token_ids}")
|
||||||
cached_response = send_single(
|
cached_response = send_single(
|
||||||
|
|||||||
Reference in New Issue
Block a user