[CI][RFC] Replace black-jupyter with ruff-format (#37210)
Co-authored-by: Alison Shao <a.shao@wustl.edu>
This commit is contained in:
co-authored by
Alison Shao
parent
2641e427be
commit
28262c20df
@@ -734,9 +734,9 @@ def _assert_files(filenames, *, exist=(), not_exist=()):
|
||||
for p in exist:
|
||||
assert any(p in f for f in filenames), f"{p} not found in {filenames}"
|
||||
for p in not_exist:
|
||||
assert not any(
|
||||
p in f for f in filenames
|
||||
), f"{p} should not exist in {filenames}"
|
||||
assert not any(p in f for f in filenames), (
|
||||
f"{p} should not exist in {filenames}"
|
||||
)
|
||||
|
||||
|
||||
def _load_dump(path: Path) -> dict:
|
||||
@@ -750,9 +750,9 @@ def _find_dump_file(tmpdir, *, rank: int = 0, name: str) -> Path:
|
||||
for f in Path(tmpdir).glob("*/*.pt")
|
||||
if f"rank={rank}" in f.name and name in f.name
|
||||
]
|
||||
assert (
|
||||
len(matches) == 1
|
||||
), f"Expected 1 file matching rank={rank} name={name}, got {matches}"
|
||||
assert len(matches) == 1, (
|
||||
f"Expected 1 file matching rank={rank} name={name}, got {matches}"
|
||||
)
|
||||
return matches[0]
|
||||
|
||||
|
||||
@@ -1657,9 +1657,9 @@ class TestZmqPortIsolation:
|
||||
)
|
||||
resp.raise_for_status()
|
||||
states = resp.json()
|
||||
assert (
|
||||
len(states) == 2
|
||||
), f"Instance {i} (port {port}): expected 2 ranks, got {len(states)}"
|
||||
assert len(states) == 2, (
|
||||
f"Instance {i} (port {port}): expected 2 ranks, got {len(states)}"
|
||||
)
|
||||
finally:
|
||||
for event in stop_events:
|
||||
event.set()
|
||||
@@ -1719,9 +1719,9 @@ class TestDumperHttp:
|
||||
val = state
|
||||
for k in keys:
|
||||
val = val[k]
|
||||
assert (
|
||||
val == expected
|
||||
), f"rank {rank}: {path}={val!r}, expected {expected!r}"
|
||||
assert val == expected, (
|
||||
f"rank {rank}: {path}={val!r}, expected {expected!r}"
|
||||
)
|
||||
|
||||
def test_configure_enable_toggle(self, dumper_http_url: str):
|
||||
for enable in [True, False]:
|
||||
@@ -1915,9 +1915,9 @@ class TestNonIntrusiveDumper(_NonIntrusiveTestBase):
|
||||
)
|
||||
|
||||
dumped_output = captured[f"{P}model.mutator.output"]["value"]
|
||||
assert (
|
||||
dumped_output == 999.0
|
||||
).all(), "post-hook should capture outputs after forward"
|
||||
assert (dumped_output == 999.0).all(), (
|
||||
"post-hook should capture outputs after forward"
|
||||
)
|
||||
|
||||
def test_hooks_all_module_levels(self, tmp_path):
|
||||
class Attention(torch.nn.Module):
|
||||
@@ -2374,9 +2374,9 @@ class TestDumperE2E:
|
||||
states = requests.post(f"{base_url}/dumper/get_state", json={}).json()
|
||||
assert len(states) == 2
|
||||
for rank, state in enumerate(states):
|
||||
assert (
|
||||
state["config"]["enable"] is True
|
||||
), f"rank {rank}: enable should be True after configure"
|
||||
assert state["config"]["enable"] is True, (
|
||||
f"rank {rank}: enable should be True after configure"
|
||||
)
|
||||
assert state["config"]["dir"] == dump_dir
|
||||
|
||||
resp = requests.post(
|
||||
@@ -2403,16 +2403,16 @@ class TestDumperE2E:
|
||||
)
|
||||
|
||||
for rank in range(2):
|
||||
assert any(
|
||||
f"rank={rank}" in f for f in filenames
|
||||
), f"No dump files for rank {rank}"
|
||||
assert any(f"rank={rank}" in f for f in filenames), (
|
||||
f"No dump files for rank {rank}"
|
||||
)
|
||||
|
||||
sample_file = dump_files[0]
|
||||
loaded = torch.load(sample_file, map_location="cpu", weights_only=False)
|
||||
assert isinstance(loaded, dict), f"Expected dict, got {type(loaded)}"
|
||||
assert (
|
||||
"value" in loaded and "meta" in loaded
|
||||
), f"Missing value/meta keys: {loaded.keys()}"
|
||||
assert "value" in loaded and "meta" in loaded, (
|
||||
f"Missing value/meta keys: {loaded.keys()}"
|
||||
)
|
||||
assert "name" in loaded["meta"]
|
||||
assert "rank" in loaded["meta"]
|
||||
assert "step" in loaded["meta"]
|
||||
@@ -2438,22 +2438,22 @@ class TestDumperE2E:
|
||||
"attn_cp_size",
|
||||
]
|
||||
for key in expected_keys:
|
||||
assert (
|
||||
key in par
|
||||
), f"Missing {key} in sglang_parallel_info, got: {sorted(par)}"
|
||||
assert key in par, (
|
||||
f"Missing {key} in sglang_parallel_info, got: {sorted(par)}"
|
||||
)
|
||||
|
||||
rids_files = [f for f in dump_files if "name=rids" in f.name]
|
||||
rids_loaded = torch.load(
|
||||
rids_files[0], map_location="cpu", weights_only=False
|
||||
)
|
||||
rids_value = rids_loaded["value"]
|
||||
assert isinstance(
|
||||
rids_value, list
|
||||
), f"rids should be a list, got {type(rids_value)}"
|
||||
assert isinstance(rids_value, list), (
|
||||
f"rids should be a list, got {type(rids_value)}"
|
||||
)
|
||||
assert len(rids_value) > 0, "rids should be non-empty"
|
||||
assert all(
|
||||
isinstance(r, str) for r in rids_value
|
||||
), f"each rid should be a str, got {[type(r) for r in rids_value]}"
|
||||
assert all(isinstance(r, str) for r in rids_value), (
|
||||
f"each rid should be a str, got {[type(r) for r in rids_value]}"
|
||||
)
|
||||
finally:
|
||||
kill_process_tree(proc.pid)
|
||||
|
||||
@@ -2912,9 +2912,9 @@ class TestRecomputeStatus:
|
||||
model(torch.randn(2, 4))
|
||||
|
||||
for key, data in captured.items():
|
||||
assert (
|
||||
"recompute_status" in data["meta"]
|
||||
), f"missing recompute_status in {key}"
|
||||
assert "recompute_status" in data["meta"], (
|
||||
f"missing recompute_status in {key}"
|
||||
)
|
||||
assert data["meta"]["recompute_status"] == "disabled"
|
||||
|
||||
def test_detect_recompute_status_default(self) -> None:
|
||||
@@ -3553,8 +3553,7 @@ class TestGrafterDistributed:
|
||||
# worker prepends tmp_path to sys.path so import_module sees it.
|
||||
module_name = "_xform_user_basic"
|
||||
(tmp_path / f"{module_name}.py").write_text(
|
||||
"def transform(graft_input):\n"
|
||||
" return graft_input.received_list[0] * 2\n"
|
||||
"def transform(graft_input):\n return graft_input.received_list[0] * 2\n"
|
||||
)
|
||||
graft_port = find_available_port(29610)
|
||||
_run_graft_test(
|
||||
@@ -3646,7 +3645,9 @@ class TestGrafterDistributed:
|
||||
7.0,
|
||||
7.0,
|
||||
7.0,
|
||||
], f"target should be unchanged after shape-mismatch graft, got {target.tolist()}"
|
||||
], (
|
||||
f"target should be unchanged after shape-mismatch graft, got {target.tolist()}"
|
||||
)
|
||||
finally:
|
||||
if grafter._pg is not None:
|
||||
dist.destroy_process_group(grafter._pg)
|
||||
@@ -3694,7 +3695,9 @@ class TestGrafterDistributed:
|
||||
9.0,
|
||||
9.0,
|
||||
9.0,
|
||||
], f"target must be unchanged when transform throws, got {target.tolist()}"
|
||||
], (
|
||||
f"target must be unchanged when transform throws, got {target.tolist()}"
|
||||
)
|
||||
output = captured.getvalue()
|
||||
assert "transform/copy_ raised RuntimeError" in output, output
|
||||
assert "intentional test error" in output, output
|
||||
@@ -3783,9 +3786,9 @@ class TestGrafterDistributed:
|
||||
grafter.maybe_intercept(value=target, tags={"name": "x"})
|
||||
output = captured.getvalue()
|
||||
if rank == 0:
|
||||
assert (
|
||||
"WARNING" in output
|
||||
), f"expected WARNING in rank 0 output: {output}"
|
||||
assert "WARNING" in output, (
|
||||
f"expected WARNING in rank 0 output: {output}"
|
||||
)
|
||||
assert "has not completed after 2s" in output, output
|
||||
finally:
|
||||
if grafter._pg is not None:
|
||||
@@ -3852,9 +3855,9 @@ class TestGrafterDistributed:
|
||||
pg_after_first = grafter._pg
|
||||
assert pg_after_first is not None
|
||||
grafter.maybe_intercept(value=t2, tags={"name": "x"})
|
||||
assert (
|
||||
grafter._pg is pg_after_first
|
||||
), "_pg must be cached across calls, not re-initialized"
|
||||
assert grafter._pg is pg_after_first, (
|
||||
"_pg must be cached across calls, not re-initialized"
|
||||
)
|
||||
else:
|
||||
target1 = torch.zeros(3, device="cuda:1")
|
||||
target2 = torch.zeros(3, device="cuda:1")
|
||||
@@ -4190,9 +4193,9 @@ def _e2e_transform(graft_input):
|
||||
the transform is just identity. Real workflows would compute a
|
||||
non-trivial override (scale, reshape, decode, ...) using the extras.
|
||||
"""
|
||||
assert (
|
||||
graft_input.received_extras_list[0]["my_extra_key"] == "my_extra_value"
|
||||
), graft_input.received_extras_list
|
||||
assert graft_input.received_extras_list[0]["my_extra_key"] == "my_extra_value", (
|
||||
graft_input.received_extras_list
|
||||
)
|
||||
return graft_input.received_list[0]
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user