Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 17 additions & 6 deletions demo/gpt2_generate.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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:
Expand All@@ -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()
Expand All@@ -45,16 +58,14 @@ 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

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")
Expand Down
3 changes: 1 addition & 2 deletions demo/warpforth.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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()),
Expand Down
15 changes: 8 additions & 7 deletions gpu_test/conftest.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -552,10 +552,11 @@ def _parse_kernel_name(forth_source: str) -> str:
"""Parse '\\! kernel <name>' 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 <name>'"
raise ValueError(msg)
return parts[1]
raise ValueError(msg) from None
msg = "Forth source has no '\\! kernel' declaration"
raise ValueError(msg)

Expand All@@ -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 <name> <type>'"
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))
Expand Down
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line numberDiff line numberDiff line change
Expand Up@@ -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"]
Loading