diff --git a/demo/gpt2_generate.py b/demo/gpt2_generate.py index 05b1838..76d7327 100644 --- a/demo/gpt2_generate.py +++ b/demo/gpt2_generate.py @@ -15,14 +15,20 @@ from __future__ import annotations import argparse +from typing import TYPE_CHECKING import torch import transformers.models.gpt2.modeling_gpt2 as gpt2_module from transformers import GPT2LMHeadModel, GPT2Tokenizer from warpforth import AttentionKernel +if TYPE_CHECKING: + from collections.abc import Callable -def make_warpforth_eager_attn(attn_kernel: AttentionKernel): + +def make_warpforth_eager_attn( + attn_kernel: AttentionKernel, +) -> Callable[..., tuple[torch.Tensor, None]]: """Create a replacement for eager_attention_forward using the WarpForth kernel. The transformers eager_attention_forward signature is: @@ -32,8 +38,15 @@ def make_warpforth_eager_attn(attn_kernel: AttentionKernel): (batch, seq_len, n_heads, head_dim). """ - def warpforth_eager_attn(module, query, key, value, attention_mask=None, **kwargs): - _batch, n_heads, seq_len, head_dim = query.shape + def warpforth_eager_attn( + _module: object, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + _attention_mask: torch.Tensor | None = None, + **_kwargs: object, + ) -> tuple[torch.Tensor, None]: + _batch, n_heads, _seq_len, _head_dim = query.shape query = query.contiguous() key = key.contiguous() value = value.contiguous() @@ -45,8 +58,6 @@ def warpforth_eager_attn(module, query, key, value, attention_mask=None, **kwarg key[0, h], value[0, h], attn_output[0, h], - seq_len, - head_dim, ) return attn_output.transpose(1, 2), None @@ -54,7 +65,7 @@ def warpforth_eager_attn(module, query, key, value, attention_mask=None, **kwarg return warpforth_eager_attn -def main(): +def main() -> None: parser = argparse.ArgumentParser(description="GPT-2 generation with WarpForth attention") parser.add_argument("--ptx", required=True, help="Path to compiled attention.ptx") parser.add_argument("--prompt", default="The meaning of life is", help="Input prompt") diff --git a/demo/warpforth.py b/demo/warpforth.py index c72e479..c9de20c 100644 --- a/demo/warpforth.py +++ b/demo/warpforth.py @@ -34,9 +34,8 @@ def __call__( k: object, # torch.Tensor (seq_len, head_dim) float32 CUDA v: object, # torch.Tensor (seq_len, head_dim) float32 CUDA o: object, # torch.Tensor (seq_len, head_dim) float32 CUDA - seq_len: int, - head_dim: int, ) -> None: + seq_len, head_dim = q.shape self._function( np.intp(q.data_ptr()), np.intp(k.data_ptr()), diff --git a/gpu_test/conftest.py b/gpu_test/conftest.py index f7fdf6a..8578f5b 100644 --- a/gpu_test/conftest.py +++ b/gpu_test/conftest.py @@ -552,10 +552,11 @@ def _parse_kernel_name(forth_source: str) -> str: """Parse '\\! kernel ' from Forth source header.""" for keyword, parts in _iter_header_directives(forth_source): if keyword == "kernel": - if len(parts) < 2: + try: + return parts[1] + except IndexError: msg = "Invalid header line: expected '\\! kernel '" - raise ValueError(msg) - return parts[1] + raise ValueError(msg) from None msg = "Forth source has no '\\! kernel' declaration" raise ValueError(msg) @@ -570,11 +571,11 @@ def _parse_param_declarations(forth_source: str) -> list[ParamDecl]: for keyword, parts in _iter_header_directives(forth_source): if keyword != "param": continue - if len(parts) < 3: + try: + name, type_spec = parts[1:3] + except ValueError: msg = "Invalid header line: expected '\\! param '" - raise ValueError(msg) - name = parts[1] - type_spec = parts[2] + raise ValueError(msg) from None if "[" in type_spec: base_type, size = _parse_array_type(type_spec) decls.append(ParamDecl(name=name, is_array=True, size=size, base_type=base_type)) diff --git a/pyproject.toml b/pyproject.toml index f416b7e..b1c1386 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -36,5 +36,5 @@ ignore = ["D", "COM812", "ISC001"] [tool.ruff.lint.per-file-ignores] "gpu_test/test_*.py" = ["S101", "PLR2004"] "gpu_test/test_vast_session.py" = ["SLF001"] -"gpu_test/conftest.py" = ["S603", "S607", "PLR0913", "PLR2004"] -"demo/*.py" = ["T201", "INP001", "SLF001", "ANN", "ARG001", "PLR0913", "PLR2004"] +"gpu_test/conftest.py" = ["S603", "S607", "PLR0913"] +"demo/*.py" = ["T201", "INP001"]