[Scheduler] Unify idle checks into is_fully_idle() and fix weight update test (#20296)

This commit is contained in:
Liangsheng Yin
2026-03-10 17:50:23 -07:00
committed by GitHub
parent 7a1ca53805
commit 50953aea8d
4 changed files with 72 additions and 105 deletions
@@ -249,44 +249,38 @@ class TestServerUpdateWeightsFromTensorNonBlocking(CustomTestCase):
return ret
def test_update_weights(self):
pause_generation_modes = ["in_place", "retract"]
for pause_generation_mode in pause_generation_modes:
num_requests = 32
with ThreadPoolExecutor(num_requests) as executor:
futures = [
executor.submit(self.run_decode, 3000) for _ in range(num_requests)
]
num_requests = 32
with ThreadPoolExecutor(num_requests) as executor:
futures = [
executor.submit(self.run_decode, 3000) for _ in range(num_requests)
]
# ensure the decode has been started
time.sleep(2)
# ensure the decode has been started
time.sleep(2)
param_names = [
f"model.layers.{i}.mlp.up_proj.weight" for i in range(6, 16)
]
new_tensor = torch.full((16384, 2048), 1.5, device="cuda")
named_tensors = [(x, new_tensor) for x in param_names]
param_names = [f"model.layers.{i}.mlp.up_proj.weight" for i in range(6, 16)]
new_tensor = torch.full((16384, 2048), 1.5, device="cuda")
named_tensors = [(x, new_tensor) for x in param_names]
ret = self.pause_generation(pause_generation_mode)
ret = self.run_update_weights(
named_tensors, flush_cache=pause_generation_mode == "retract"
# abort mode ensures server is totally idle before returning
ret = self.pause_generation("abort")
ret = self.run_update_weights(named_tensors, flush_cache=True)
self.assertTrue(ret["success"])
ret = self.continue_generation()
# requests were aborted by pause_generation("abort")
for future in as_completed(futures):
future.result()
for param_name in param_names[:3]:
response = requests.post(
self.base_url + "/get_weights_by_name",
json={"name": param_name},
)
self.assertTrue(ret["success"])
ret = self.continue_generation()
for future in as_completed(futures):
self.assertNotEqual(
future.result()["meta_info"]["finish_reason"]["type"], "abort"
)
for param_name in param_names[:3]:
response = requests.post(
self.base_url + "/get_weights_by_name",
json={"name": param_name},
)
actual_values = torch.tensor(response.json())[0, :5]
assert torch.allclose(
actual_values, torch.tensor([1.5] * 5), atol=0.002
), f"{actual_values=}"
actual_values = torch.tensor(response.json())[0, :5]
assert torch.allclose(
actual_values, torch.tensor([1.5] * 5), atol=0.002
), f"{actual_values=}"
def _check_param(engine, param_name, expect_values):