From 1b2bd65e3861bb19f04ab7ff7ab55150ebcba4c9 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Tue, 21 Apr 2026 04:58:24 +0000 Subject: [PATCH] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- scripts/benchmarks/flex_paged_attention.py | 2 +- scripts/benchmarks/qwen3_flex_inference.py | 25 +++++++++++----------- 2 files changed, 14 insertions(+), 13 deletions(-) diff --git a/scripts/benchmarks/flex_paged_attention.py b/scripts/benchmarks/flex_paged_attention.py index 146855675b..e76716ab46 100644 --- a/scripts/benchmarks/flex_paged_attention.py +++ b/scripts/benchmarks/flex_paged_attention.py @@ -17,7 +17,7 @@ from torch.nn.attention.flex_attention import ( create_block_mask, ) -create_block_mask = torch.compile(create_block_mask, dynamic=True) +create_block_mask = torch.compile(create_block_mask, dynamic = True) def _cdiv(x: int | float | torch.Tensor, multiple: int | float | torch.Tensor): diff --git a/scripts/benchmarks/qwen3_flex_inference.py b/scripts/benchmarks/qwen3_flex_inference.py index e78c63c120..d94f695d4b 100644 --- a/scripts/benchmarks/qwen3_flex_inference.py +++ b/scripts/benchmarks/qwen3_flex_inference.py @@ -308,8 +308,8 @@ def refresh_lora_merge_from_pristine(base_model, peft_model): pristine_w, B.to(W.dtype), A.to(W.dtype), - alpha=module.scaling[adapter0], - out=W, + alpha = module.scaling[adapter0], + out = W, ) for adapter in active[1:]: A = module.lora_A[adapter].weight.data @@ -318,8 +318,8 @@ def refresh_lora_merge_from_pristine(base_model, peft_model): W, B.to(W.dtype), A.to(W.dtype), - alpha=module.scaling[adapter], - out=W, + alpha = module.scaling[adapter], + out = W, ) module.merged_adapters = list(active) n_refreshed += 1 @@ -343,8 +343,9 @@ def _hash_state_dict(model) -> str: return h.hexdigest() -def run_drift_verification(base_model, peft_model, n_iters: int = 10, - noise_scale: float = 0.01): +def run_drift_verification( + base_model, peft_model, n_iters: int = 10, noise_scale: float = 0.01 +): """Simulate N GRPO iterations: perturb LoRA weights with random noise, call `refresh_lora_merge_from_pristine`, repeat. Assert the pristine `base_model`'s parameters are bit-identical before and after. @@ -363,12 +364,12 @@ def run_drift_verification(base_model, peft_model, n_iters: int = 10, if not isinstance(module, LoraLayer): continue for adapter_name in list(module.lora_A.keys()): - initial_lora[(name, "A", adapter_name)] = ( - module.lora_A[adapter_name].weight.data.clone() - ) - initial_lora[(name, "B", adapter_name)] = ( - module.lora_B[adapter_name].weight.data.clone() - ) + initial_lora[(name, "A", adapter_name)] = module.lora_A[ + adapter_name + ].weight.data.clone() + initial_lora[(name, "B", adapter_name)] = module.lora_B[ + adapter_name + ].weight.data.clone() base_hash_before = _hash_state_dict(base_model)