liangsu9988 commited on
Commit
bbb0f58
·
verified ·
1 Parent(s): b02f9a3

Add torch211-cxx11-cu130-aarch64-linux SM110 artifact

Browse files
build/torch211-cxx11-cu130-aarch64-linux/__init__.py ADDED
@@ -0,0 +1,105 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """FlashRT speculative decoding helper kernels."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import Optional
6
+
7
+ import torch
8
+
9
+ from ._ops import add_op_namespace_prefix, ops
10
+
11
+
12
+ @torch.library.register_fake(add_op_namespace_prefix("argmax_bf16"))
13
+ def _argmax_bf16_fake(logits: torch.Tensor, argmax_out: torch.Tensor) -> None:
14
+ if logits.dim() != 2 or argmax_out.shape != (logits.shape[0],):
15
+ raise RuntimeError("argmax_bf16 expects logits (rows,vocab), argmax_out (rows,)")
16
+ return None
17
+
18
+
19
+ @torch.library.register_fake(add_op_namespace_prefix("accept_greedy_bf16"))
20
+ def _accept_greedy_bf16_fake(
21
+ logits: torch.Tensor,
22
+ drafts: torch.Tensor,
23
+ argmax_out: torch.Tensor,
24
+ accept_n: torch.Tensor,
25
+ spec_k: int,
26
+ ) -> None:
27
+ if logits.dim() != 2 or argmax_out.shape != (logits.shape[0],):
28
+ raise RuntimeError("accept_greedy_bf16 expects logits (rows,vocab), argmax_out (rows,)")
29
+ if drafts.dim() != 1 or drafts.numel() < spec_k or accept_n.numel() < 1:
30
+ raise RuntimeError("drafts/accept_n shape mismatch")
31
+ return None
32
+
33
+
34
+ @torch.library.register_fake(add_op_namespace_prefix("accept_partitioned_bf16"))
35
+ def _accept_partitioned_bf16_fake(
36
+ logits: torch.Tensor,
37
+ drafts: torch.Tensor,
38
+ argmax_out: torch.Tensor,
39
+ accept_n: torch.Tensor,
40
+ partial_vals: torch.Tensor,
41
+ partial_idx: torch.Tensor,
42
+ spec_k: int,
43
+ parts: int,
44
+ ) -> None:
45
+ if partial_vals.shape != (logits.shape[0], parts) or partial_idx.shape != (logits.shape[0], parts):
46
+ raise RuntimeError("partial buffers must have shape (rows, parts)")
47
+ return _accept_greedy_bf16_fake(logits, drafts, argmax_out, accept_n, spec_k)
48
+
49
+
50
+ def argmax_bf16(logits: torch.Tensor, *, out: Optional[torch.Tensor] = None) -> torch.Tensor:
51
+ if out is None:
52
+ out = torch.empty((logits.shape[0],), device=logits.device, dtype=torch.int64)
53
+ ops.argmax_bf16(logits, out)
54
+ return out
55
+
56
+
57
+ def accept_greedy_bf16(
58
+ logits: torch.Tensor,
59
+ drafts: torch.Tensor,
60
+ spec_k: int,
61
+ *,
62
+ argmax_out: Optional[torch.Tensor] = None,
63
+ accept_n: Optional[torch.Tensor] = None,
64
+ ) -> tuple[torch.Tensor, torch.Tensor]:
65
+ if argmax_out is None:
66
+ argmax_out = torch.empty((logits.shape[0],), device=logits.device, dtype=torch.int64)
67
+ if accept_n is None:
68
+ accept_n = torch.empty((1,), device=logits.device, dtype=torch.int32)
69
+ ops.accept_greedy_bf16(logits, drafts, argmax_out, accept_n, int(spec_k))
70
+ return argmax_out, accept_n
71
+
72
+
73
+ def accept_partitioned_bf16(
74
+ logits: torch.Tensor,
75
+ drafts: torch.Tensor,
76
+ spec_k: int,
77
+ parts: Optional[int] = None,
78
+ *,
79
+ argmax_out: Optional[torch.Tensor] = None,
80
+ accept_n: Optional[torch.Tensor] = None,
81
+ partial_vals: Optional[torch.Tensor] = None,
82
+ partial_idx: Optional[torch.Tensor] = None,
83
+ ) -> tuple[torch.Tensor, torch.Tensor]:
84
+ if parts is None:
85
+ vocab = int(logits.shape[1])
86
+ parts = 32 if vocab >= 131072 else (16 if vocab >= 65536 else 8)
87
+ if argmax_out is None:
88
+ argmax_out = torch.empty((logits.shape[0],), device=logits.device, dtype=torch.int64)
89
+ if accept_n is None:
90
+ accept_n = torch.empty((1,), device=logits.device, dtype=torch.int32)
91
+ if partial_vals is None:
92
+ partial_vals = torch.empty((logits.shape[0], parts), device=logits.device, dtype=torch.float32)
93
+ if partial_idx is None:
94
+ partial_idx = torch.empty((logits.shape[0], parts), device=logits.device, dtype=torch.int32)
95
+ ops.accept_partitioned_bf16(
96
+ logits, drafts, argmax_out, accept_n, partial_vals, partial_idx, int(spec_k), int(parts)
97
+ )
98
+ return argmax_out, accept_n
99
+
100
+
101
+ __all__ = [
102
+ "argmax_bf16",
103
+ "accept_greedy_bf16",
104
+ "accept_partitioned_bf16",
105
+ ]
build/torch211-cxx11-cu130-aarch64-linux/_ops.py ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ import torch
2
+ from . import _speculative_draft_primitives_cuda_6ee7cab
3
+ ops = torch.ops._speculative_draft_primitives_cuda_6ee7cab
4
+
5
+ def add_op_namespace_prefix(op_name: str):
6
+ return f"_speculative_draft_primitives_cuda_6ee7cab::{op_name}"
build/torch211-cxx11-cu130-aarch64-linux/_speculative_draft_primitives_cuda_6ee7cab.abi3.so ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6b95428d63268aaf762433b0bbf183219fa0571811edb8c31ea483233f41a60a
3
+ size 175736
build/torch211-cxx11-cu130-aarch64-linux/metadata.json ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "name": "speculative-draft-primitives",
3
+ "id": "_speculative_draft_primitives_cuda_6ee7cab",
4
+ "version": 1,
5
+ "license": "Apache-2.0",
6
+ "python-depends": [],
7
+ "backend": {
8
+ "type": "cuda",
9
+ "archs": [
10
+ "11.0"
11
+ ]
12
+ },
13
+ "digest": {
14
+ "algorithm": "sha256",
15
+ "files": {
16
+ "__init__.py": "6NhsYcUvz+6W9eEWWCQJVI8aaXA89WhNaHArHAM/YiE=",
17
+ "_speculative_draft_primitives_cuda_6ee7cab.abi3.so": "a5VCjWMmiq92JDOwu/GDIZ+gVxgR7bjDHqSDIz9Bpgo=",
18
+ "_ops.py": "mXBQn2p+IHGD95tP5Z0LrYSEmTCuehbwhHIOcFElJg0=",
19
+ "speculative_draft_primitives/__init__.py": "v6p5XMfQzddhi1fLSAw4HX9CyS0rQsidvu9VsT01xi4="
20
+ }
21
+ }
22
+ }
build/torch211-cxx11-cu130-aarch64-linux/speculative_draft_primitives/__init__.py ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import ctypes
2
+ import importlib.util
3
+ import sys
4
+ from pathlib import Path
5
+
6
+ def _import_from_path(file_path: Path):
7
+ path_hash = '{:x}'.format(ctypes.c_size_t(hash(file_path.absolute())).value)
8
+ spec = importlib.util.spec_from_file_location(path_hash, file_path)
9
+ module = importlib.util.module_from_spec(spec)
10
+ sys.modules[path_hash] = module
11
+ spec.loader.exec_module(module)
12
+ return module
13
+
14
+ globals().update(vars(_import_from_path(Path(__file__).parent.parent / '__init__.py')))