(self)
| 523 | self._applied: List[str] = [] |
| 524 | |
| 525 | def __enter__(self) -> nn.Module: |
| 526 | for repl in self.replacements: |
| 527 | try: |
| 528 | kernel_mod = load_kernel_module(repl.optimized_path) |
| 529 | if not hasattr(kernel_mod, "kernel_fn"): |
| 530 | print(f" WARNING: {repl.optimized_path} has no kernel_fn, skipping") |
| 531 | continue |
| 532 | repl.module_fn = kernel_mod.kernel_fn |
| 533 | except Exception as e: |
| 534 | print(f" WARNING: Failed to load {repl.optimized_path}: {e}") |
| 535 | continue |
| 536 | |
| 537 | replaced = self._apply_replacement(repl) |
| 538 | if replaced > 0: |
| 539 | self._applied.append( |
| 540 | f" {repl.kernel_type} (rank {repl.rank}): " |
| 541 | f"{repl.speedup:.1f}x -> {repl.optimized_path}" |
| 542 | ) |
| 543 | |
| 544 | return self.model |
| 545 | |
| 546 | def __exit__(self, *exc): |
| 547 | # Restore all original modules |
nothing calls this directly
no test coverage detected