| |
| |
| |
| |
| |
| |
| |
| |
|
|
| from __future__ import annotations |
|
|
| from pathlib import Path |
|
|
| import kernels |
| import torch |
| import torch.nn.functional as F |
|
|
|
|
| def main() -> None: |
| project = Path(__file__).parent |
| build = project / "result" if (project / "result").exists() else project / "build" |
| linear_attention = kernels.get_local_kernel(build, "cuda") |
|
|
| b, t, h, d = 2, 512, 8, 64 |
| q = torch.randn(b, t, h, d, device="cuda", dtype=torch.bfloat16) |
| k = F.normalize( |
| torch.randn(b, t, h, d, device="cuda", dtype=torch.float32), dim=-1 |
| ).to(q) |
| v = torch.randn_like(q) |
| g = -torch.rand(b, t, h, d, device="cuda", dtype=torch.float32) * 0.1 |
| beta = torch.rand(b, t, h, device="cuda", dtype=torch.bfloat16) |
|
|
| output, final_state = linear_attention.chunk_kda( |
| q, |
| k, |
| v, |
| g, |
| beta, |
| output_final_state=True, |
| ) |
| print("output:", tuple(output.shape)) |
| print("final state:", tuple(final_state.shape)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|