Context manager that patches a model's submodules to use optimized Triton kernels. Usage: with OptimizedModelContext(model, replacements) as patched_model: output = patched_model(input)
| 508 | |
| 509 | |
| 510 | class OptimizedModelContext: |
| 511 | """ |
| 512 | Context manager that patches a model's submodules to use optimized Triton kernels. |
| 513 | |
| 514 | Usage: |
| 515 | with OptimizedModelContext(model, replacements) as patched_model: |
| 516 | output = patched_model(input) |
| 517 | """ |
| 518 | |
| 519 | def __init__(self, model: nn.Module, replacements: List[KernelReplacement]): |
| 520 | self.model = model |
| 521 | self.replacements = replacements |
| 522 | self._original_modules: Dict[str, nn.Module] = {} |
| 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 |
| 548 | for name, original in self._original_modules.items(): |
| 549 | parts = name.split(".") |
| 550 | parent = self.model |
| 551 | for p in parts[:-1]: |
| 552 | parent = getattr(parent, p) |
| 553 | setattr(parent, parts[-1], original) |
| 554 | self._original_modules.clear() |
| 555 | self._applied.clear() |
| 556 | |
| 557 | def _apply_replacement(self, repl: KernelReplacement) -> int: |
| 558 | """ |
| 559 | Replace matching modules in the model. Returns number of modules replaced. |
| 560 | """ |
| 561 | count = 0 |
| 562 | |
| 563 | if repl.kernel_type == "matmul": |
| 564 | count = self._replace_linear_modules(repl) |
| 565 | elif repl.kernel_type == "layernorm": |
| 566 | count = self._replace_layernorm_modules(repl) |
| 567 | elif repl.kernel_type == "rmsnorm": |
no outgoing calls
no test coverage detected