From 8daa9577d94e194d8b6be29edcc9281cee00a496 Mon Sep 17 00:00:00 2001 From: cenzhiyao <2523403608@qq.com> Date: Sat, 26 Sep 2026 04:35:41 +0000 Subject: [PATCH] fix(test): increase model dim for stable timing benchmark MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Bump SimpleModel dim 128→1024 and seq_len 256→512 so GPU compute dominates over kernel-launch overhead, eliminating flaky ratio assertions on fast GPUs (B300/H100). --- tests/api_tests/test_magi_compile.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/tests/api_tests/test_magi_compile.py b/tests/api_tests/test_magi_compile.py index 8b37345..65957b8 100644 --- a/tests/api_tests/test_magi_compile.py +++ b/tests/api_tests/test_magi_compile.py @@ -552,7 +552,7 @@ def test_simple_model_timing_class_function_instance_method(self): class SimpleModel(nn.Module): def __init__(self): super().__init__() - self.dim = 128 + self.dim = 1024 self.layers = nn.ModuleList([nn.Linear(self.dim, self.dim) for _ in range(4)]) self.norms = nn.ModuleList([nn.LayerNorm(self.dim) for _ in range(4)]) @@ -567,8 +567,8 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: return x device = torch.device("cuda:0") - seq_len = 256 - test_input = torch.randn(seq_len, 128, device=device) + seq_len = 512 + test_input = torch.randn(seq_len, 1024, device=device) native = SimpleModel().to(device).eval() @@ -632,16 +632,16 @@ class CompiledNonModulePerf(NonModulePerf): def forward(self, x: torch.Tensor) -> torch.Tensor: return super().forward(x) - non_module_native = NonModulePerf(128) + non_module_native = NonModulePerf(1024) - non_module_cls = CompiledNonModulePerf(128) + non_module_cls = CompiledNonModulePerf(1024) non_module_cls.copy_from(non_module_native) - non_module_inst_obj = NonModulePerf(128) + non_module_inst_obj = NonModulePerf(1024) non_module_inst_obj.copy_from(non_module_native) non_module_inst = magi_compile(non_module_inst_obj, dynamic_arg_dims={"x": 0}) - non_module_mtd_obj = NonModulePerf(128) + non_module_mtd_obj = NonModulePerf(1024) non_module_mtd_obj.copy_from(non_module_native) non_module_mtd_obj.step = magi_compile(non_module_mtd_obj.step, dynamic_arg_dims={"x": 0})