fix(io_struct): index extra_key per sub-request in batched GenerateReqInput (#26971)

This commit is contained in:
Humphrey
2026-06-14 00:50:38 -07:00
committed by GitHub
parent 1456eb612d
commit 8c334e2224
2 changed files with 68 additions and 1 deletions
@@ -419,6 +419,57 @@ class TestGenerateReqInputNormalization(CustomTestCase):
req.normalize_batch_and_arguments()
self.assertEqual(req.lora_path, expected_lora_paths)
def test_extra_key_normalization(self):
"""Test normalization of extra_key."""
# Per-request list
req = GenerateReqInput(
text=["Hello", "World"],
extra_key=["tenant-A", "tenant-B"],
sampling_params=[{}, {}],
)
req.normalize_batch_and_arguments()
self.assertEqual(req.extra_key, ["tenant-A", "tenant-B"])
self.assertEqual(req[0].extra_key, "tenant-A")
self.assertEqual(req[1].extra_key, "tenant-B")
# Scalar broadcast
req = GenerateReqInput(
text=["Hello", "World"],
extra_key="shared",
sampling_params=[{}, {}],
)
req.normalize_batch_and_arguments()
self.assertEqual(req.extra_key, ["shared", "shared"])
# None stays None
req = GenerateReqInput(text=["Hello", "World"], sampling_params=[{}, {}])
req.normalize_batch_and_arguments()
self.assertIsNone(req.extra_key)
self.assertIsNone(req[0].extra_key)
# Parallel sampling expansion
req = GenerateReqInput(
text=["Hello", "World"],
extra_key=["tenant-A", "tenant-B"],
sampling_params={"n": 2},
)
req.normalize_batch_and_arguments()
self.assertEqual(req.extra_key, ["tenant-A", "tenant-B"] * 2)
# Wrong-length list
req = GenerateReqInput(
text=["Hello", "World"],
extra_key=["only-one"],
sampling_params=[{}, {}],
)
with self.assertRaisesRegex(ValueError, "batch size"):
req.normalize_batch_and_arguments()
# Non-batched scalar unchanged
req = GenerateReqInput(text="Hello", extra_key="solo")
req.normalize_batch_and_arguments()
self.assertEqual(req.extra_key, "solo")
def test_logprob_parameters_normalization(self):
"""Test normalization of logprob-related parameters."""
# Test single example