liangsu9988 commited on
Commit
aa2ae30
·
verified ·
1 Parent(s): 59480de

Uploaded using `kernel-builder`.

Browse files
benchmarks/benchmark.py CHANGED
@@ -215,7 +215,8 @@ def main() -> int:
215
  "small_m16_n128_k128": (16, 128, 128),
216
  "small_m32_n256_k256": (32, 256, 256),
217
  "mlp_tile_m64_n512_k512": (64, 512, 512),
218
- "groot_dit_projection": (51, 1536, 1536),
 
219
  "vla_projection": (105, 2048, 2048),
220
  "motus_up": (360, 14336, 3072),
221
  "motus_down": (360, 3072, 14336),
 
215
  "small_m16_n128_k128": (16, 128, 128),
216
  "small_m32_n256_k256": (32, 256, 256),
217
  "mlp_tile_m64_n512_k512": (64, 512, 512),
218
+ "groot_n17_dit_projection": (41, 1536, 1536),
219
+ "groot_legacy_dit_projection": (51, 1536, 1536),
220
  "vla_projection": (105, 2048, 2048),
221
  "motus_up": (360, 14336, 3072),
222
  "motus_down": (360, 3072, 14336),
build/torch212-cxx11-cu130-x86_64-linux/__init__.py CHANGED
@@ -58,6 +58,16 @@ def _legacy_linear_fake(
58
  return None
59
 
60
 
 
 
 
 
 
 
 
 
 
 
61
  @torch.library.register_fake(add_op_namespace_prefix("quantize_fp4_sfa_fp16"))
62
  def _quant_fake(x: torch.Tensor, packed: torch.Tensor, sfa: torch.Tensor, is_sfb: bool = False) -> None:
63
  return None
@@ -191,6 +201,45 @@ def fp4_w4a16_linear_bf16(
191
  )
192
 
193
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
194
  def nvfp4_gemm_residual_bf16(
195
  a_packed: torch.Tensor,
196
  b_packed: torch.Tensor,
@@ -296,8 +345,10 @@ __all__ = [
296
  "fp4_w4a16_linear_bf16",
297
  "fp4_w4a4_gemv_warpsplit_bf16",
298
  "nvfp4_gemm_bf16",
 
299
  "nvfp4_gemm_bias_gelu_bf16",
300
  "nvfp4_gemm_bias_gelu_nvfp4",
 
301
  "nvfp4_gemm_residual_bf16",
302
  "nvfp4_gemm_streamk_bf16",
303
  "nvfp4_gemm_streamk_bias_bf16",
 
58
  return None
59
 
60
 
61
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_bf16"))
62
+ def _bias_fake(a, b, sfa, sfb, bias, out) -> None:
63
+ return None
64
+
65
+
66
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_residual_bf16"))
67
+ def _bias_residual_fake(a, b, sfa, sfb, bias, residual, out) -> None:
68
+ return None
69
+
70
+
71
  @torch.library.register_fake(add_op_namespace_prefix("quantize_fp4_sfa_fp16"))
72
  def _quant_fake(x: torch.Tensor, packed: torch.Tensor, sfa: torch.Tensor, is_sfb: bool = False) -> None:
73
  return None
 
201
  )
202
 
203
 
204
+ def nvfp4_gemm_bias_bf16(
205
+ a_packed: torch.Tensor,
206
+ b_packed: torch.Tensor,
207
+ sfa: torch.Tensor,
208
+ sfb: torch.Tensor,
209
+ bias: torch.Tensor,
210
+ *,
211
+ out: torch.Tensor | None = None,
212
+ ) -> torch.Tensor:
213
+ """SM110 NVFP4 GEMM with a fused per-column BF16 bias."""
214
+ if out is None:
215
+ out = torch.empty(
216
+ (a_packed.shape[0], b_packed.shape[0]),
217
+ device=a_packed.device,
218
+ dtype=torch.bfloat16,
219
+ )
220
+ ops.nvfp4_gemm_bias_bf16(a_packed, b_packed, sfa, sfb, bias, out)
221
+ return out
222
+
223
+
224
+ def nvfp4_gemm_bias_residual_bf16(
225
+ a_packed: torch.Tensor,
226
+ b_packed: torch.Tensor,
227
+ sfa: torch.Tensor,
228
+ sfb: torch.Tensor,
229
+ bias: torch.Tensor,
230
+ residual: torch.Tensor,
231
+ *,
232
+ out: torch.Tensor | None = None,
233
+ ) -> torch.Tensor:
234
+ """SM110 NVFP4 GEMM with fused BF16 bias and residual add."""
235
+ if out is None:
236
+ out = torch.empty_like(residual)
237
+ ops.nvfp4_gemm_bias_residual_bf16(
238
+ a_packed, b_packed, sfa, sfb, bias, residual, out
239
+ )
240
+ return out
241
+
242
+
243
  def nvfp4_gemm_residual_bf16(
244
  a_packed: torch.Tensor,
245
  b_packed: torch.Tensor,
 
345
  "fp4_w4a16_linear_bf16",
346
  "fp4_w4a4_gemv_warpsplit_bf16",
347
  "nvfp4_gemm_bf16",
348
+ "nvfp4_gemm_bias_bf16",
349
  "nvfp4_gemm_bias_gelu_bf16",
350
  "nvfp4_gemm_bias_gelu_nvfp4",
351
+ "nvfp4_gemm_bias_residual_bf16",
352
  "nvfp4_gemm_residual_bf16",
353
  "nvfp4_gemm_streamk_bf16",
354
  "nvfp4_gemm_streamk_bias_bf16",
build/torch212-cxx11-cu130-x86_64-linux/{_fp4_gemm_cuda_0fe564a.abi3.so → _fp4_gemm_cuda_8a66d8b.abi3.so} RENAMED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:6d34445bc1a1d2010ed592eade54e88fd2c872f135a49d18fa98d1aba9c95fe6
3
- size 2428552
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:26fce4c1404d996b8e9ab774125c787e47a568f79da35f734949bff0177a173a
3
+ size 2848128
build/torch212-cxx11-cu130-x86_64-linux/_ops.py CHANGED
@@ -1,9 +1,9 @@
1
  import torch
2
- from . import _fp4_gemm_cuda_0fe564a
3
- ops = torch.ops._fp4_gemm_cuda_0fe564a
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
- return f"_fp4_gemm_cuda_0fe564a::{op_name}"
 
1
  import torch
2
+ from . import _fp4_gemm_cuda_8a66d8b
3
+ ops = torch.ops._fp4_gemm_cuda_8a66d8b
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
+ return f"_fp4_gemm_cuda_8a66d8b::{op_name}"
build/torch212-cxx11-cu130-x86_64-linux/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "fp4-gemm",
3
- "id": "_fp4_gemm_cuda_0fe564a",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
@@ -14,19 +14,19 @@
14
  "digest": {
15
  "algorithm": "sha256",
16
  "files": {
17
- "__init__.py": "pJlGLDlOKnju+m59WyZJjpynMTvkPQTF6w+z4YfCwSY=",
18
- "_fp4_gemm_cuda_0fe564a.abi3.so": "bTREW8Gh0gEO1ZLq3lToj9LIcvE1pJ0Y+pjRq6nJX+Y=",
19
- "_ops.py": "64wwcwE28ue+jCMSXh8vVgbUevYtei6uGZjpzaoMVrw="
20
  }
21
  },
22
  "provenance": {
23
  "kernel-builder": {
24
  "version": "0.17.0-dev0",
25
- "sha": "d720fa90fb9cd92d1bc60a9dc5c55bef2aafabb8",
26
  "dirty": false
27
  },
28
  "kernel": {
29
- "sha": "0fe564afa690cd8a731729a59f98be95ebee4b34",
30
  "dirty": false
31
  }
32
  }
 
1
  {
2
  "name": "fp4-gemm",
3
+ "id": "_fp4_gemm_cuda_8a66d8b",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
 
14
  "digest": {
15
  "algorithm": "sha256",
16
  "files": {
17
+ "__init__.py": "+Kk/VNnIWwe9nszQYvL2qtU3DpKxbqHVfklg5WtVuB4=",
18
+ "_fp4_gemm_cuda_8a66d8b.abi3.so": "JvzkwUBNmWuOmrd0Elx4fkelaPedo19zSUm/8Bd6Fzo=",
19
+ "_ops.py": "mZSHUC0K9o9glzvGmRSdtPppOq4bmIaxNEjibkaWTk0="
20
  }
21
  },
22
  "provenance": {
23
  "kernel-builder": {
24
  "version": "0.17.0-dev0",
25
+ "sha": "870e825d881664e39f9287a27a74ef63ff3c545e",
26
  "dirty": false
27
  },
28
  "kernel": {
29
+ "sha": "8a66d8b7f79a6fde86aa6929db2b60ab097e2d55",
30
  "dirty": false
31
  }
32
  }
build/torch212-cxx11-cu132-x86_64-linux/__init__.py CHANGED
@@ -58,6 +58,16 @@ def _legacy_linear_fake(
58
  return None
59
 
60
 
 
 
 
 
 
 
 
 
 
 
61
  @torch.library.register_fake(add_op_namespace_prefix("quantize_fp4_sfa_fp16"))
62
  def _quant_fake(x: torch.Tensor, packed: torch.Tensor, sfa: torch.Tensor, is_sfb: bool = False) -> None:
63
  return None
@@ -191,6 +201,45 @@ def fp4_w4a16_linear_bf16(
191
  )
192
 
193
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
194
  def nvfp4_gemm_residual_bf16(
195
  a_packed: torch.Tensor,
196
  b_packed: torch.Tensor,
@@ -296,8 +345,10 @@ __all__ = [
296
  "fp4_w4a16_linear_bf16",
297
  "fp4_w4a4_gemv_warpsplit_bf16",
298
  "nvfp4_gemm_bf16",
 
299
  "nvfp4_gemm_bias_gelu_bf16",
300
  "nvfp4_gemm_bias_gelu_nvfp4",
 
301
  "nvfp4_gemm_residual_bf16",
302
  "nvfp4_gemm_streamk_bf16",
303
  "nvfp4_gemm_streamk_bias_bf16",
 
58
  return None
59
 
60
 
61
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_bf16"))
62
+ def _bias_fake(a, b, sfa, sfb, bias, out) -> None:
63
+ return None
64
+
65
+
66
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_residual_bf16"))
67
+ def _bias_residual_fake(a, b, sfa, sfb, bias, residual, out) -> None:
68
+ return None
69
+
70
+
71
  @torch.library.register_fake(add_op_namespace_prefix("quantize_fp4_sfa_fp16"))
72
  def _quant_fake(x: torch.Tensor, packed: torch.Tensor, sfa: torch.Tensor, is_sfb: bool = False) -> None:
73
  return None
 
201
  )
202
 
203
 
204
+ def nvfp4_gemm_bias_bf16(
205
+ a_packed: torch.Tensor,
206
+ b_packed: torch.Tensor,
207
+ sfa: torch.Tensor,
208
+ sfb: torch.Tensor,
209
+ bias: torch.Tensor,
210
+ *,
211
+ out: torch.Tensor | None = None,
212
+ ) -> torch.Tensor:
213
+ """SM110 NVFP4 GEMM with a fused per-column BF16 bias."""
214
+ if out is None:
215
+ out = torch.empty(
216
+ (a_packed.shape[0], b_packed.shape[0]),
217
+ device=a_packed.device,
218
+ dtype=torch.bfloat16,
219
+ )
220
+ ops.nvfp4_gemm_bias_bf16(a_packed, b_packed, sfa, sfb, bias, out)
221
+ return out
222
+
223
+
224
+ def nvfp4_gemm_bias_residual_bf16(
225
+ a_packed: torch.Tensor,
226
+ b_packed: torch.Tensor,
227
+ sfa: torch.Tensor,
228
+ sfb: torch.Tensor,
229
+ bias: torch.Tensor,
230
+ residual: torch.Tensor,
231
+ *,
232
+ out: torch.Tensor | None = None,
233
+ ) -> torch.Tensor:
234
+ """SM110 NVFP4 GEMM with fused BF16 bias and residual add."""
235
+ if out is None:
236
+ out = torch.empty_like(residual)
237
+ ops.nvfp4_gemm_bias_residual_bf16(
238
+ a_packed, b_packed, sfa, sfb, bias, residual, out
239
+ )
240
+ return out
241
+
242
+
243
  def nvfp4_gemm_residual_bf16(
244
  a_packed: torch.Tensor,
245
  b_packed: torch.Tensor,
 
345
  "fp4_w4a16_linear_bf16",
346
  "fp4_w4a4_gemv_warpsplit_bf16",
347
  "nvfp4_gemm_bf16",
348
+ "nvfp4_gemm_bias_bf16",
349
  "nvfp4_gemm_bias_gelu_bf16",
350
  "nvfp4_gemm_bias_gelu_nvfp4",
351
+ "nvfp4_gemm_bias_residual_bf16",
352
  "nvfp4_gemm_residual_bf16",
353
  "nvfp4_gemm_streamk_bf16",
354
  "nvfp4_gemm_streamk_bias_bf16",
build/torch212-cxx11-cu132-x86_64-linux/{_fp4_gemm_cuda_0fe564a.abi3.so → _fp4_gemm_cuda_8a66d8b.abi3.so} RENAMED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:88e788a3e8cc8ed5dcdc13a2bc9b846639d0e116b0f17d3f0ff882af3693c0ff
3
- size 2428504
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:69c1903f5c54a54d13a0a27890d26ca0f17123dfceb463d44c1cc7185d2db1b6
3
+ size 2843976
build/torch212-cxx11-cu132-x86_64-linux/_ops.py CHANGED
@@ -1,9 +1,9 @@
1
  import torch
2
- from . import _fp4_gemm_cuda_0fe564a
3
- ops = torch.ops._fp4_gemm_cuda_0fe564a
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
- return f"_fp4_gemm_cuda_0fe564a::{op_name}"
 
1
  import torch
2
+ from . import _fp4_gemm_cuda_8a66d8b
3
+ ops = torch.ops._fp4_gemm_cuda_8a66d8b
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
+ return f"_fp4_gemm_cuda_8a66d8b::{op_name}"
build/torch212-cxx11-cu132-x86_64-linux/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "fp4-gemm",
3
- "id": "_fp4_gemm_cuda_0fe564a",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
@@ -14,19 +14,19 @@
14
  "digest": {
15
  "algorithm": "sha256",
16
  "files": {
17
- "__init__.py": "pJlGLDlOKnju+m59WyZJjpynMTvkPQTF6w+z4YfCwSY=",
18
- "_fp4_gemm_cuda_0fe564a.abi3.so": "iOeIo+jMjtXc3BOivJuEZjnQ4Raw8X0/D/iCrzaTwP8=",
19
- "_ops.py": "64wwcwE28ue+jCMSXh8vVgbUevYtei6uGZjpzaoMVrw="
20
  }
21
  },
22
  "provenance": {
23
  "kernel-builder": {
24
  "version": "0.17.0-dev0",
25
- "sha": "d720fa90fb9cd92d1bc60a9dc5c55bef2aafabb8",
26
  "dirty": false
27
  },
28
  "kernel": {
29
- "sha": "0fe564afa690cd8a731729a59f98be95ebee4b34",
30
  "dirty": false
31
  }
32
  }
 
1
  {
2
  "name": "fp4-gemm",
3
+ "id": "_fp4_gemm_cuda_8a66d8b",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
 
14
  "digest": {
15
  "algorithm": "sha256",
16
  "files": {
17
+ "__init__.py": "+Kk/VNnIWwe9nszQYvL2qtU3DpKxbqHVfklg5WtVuB4=",
18
+ "_fp4_gemm_cuda_8a66d8b.abi3.so": "acGQP1xUpU0ToKJ4kNJsoPFxI9/OtGPUTBzHGF0tsbY=",
19
+ "_ops.py": "mZSHUC0K9o9glzvGmRSdtPppOq4bmIaxNEjibkaWTk0="
20
  }
21
  },
22
  "provenance": {
23
  "kernel-builder": {
24
  "version": "0.17.0-dev0",
25
+ "sha": "870e825d881664e39f9287a27a74ef63ff3c545e",
26
  "dirty": false
27
  },
28
  "kernel": {
29
+ "sha": "8a66d8b7f79a6fde86aa6929db2b60ab097e2d55",
30
  "dirty": false
31
  }
32
  }
build/torch213-cxx11-cu130-x86_64-linux/__init__.py CHANGED
@@ -58,6 +58,16 @@ def _legacy_linear_fake(
58
  return None
59
 
60
 
 
 
 
 
 
 
 
 
 
 
61
  @torch.library.register_fake(add_op_namespace_prefix("quantize_fp4_sfa_fp16"))
62
  def _quant_fake(x: torch.Tensor, packed: torch.Tensor, sfa: torch.Tensor, is_sfb: bool = False) -> None:
63
  return None
@@ -191,6 +201,45 @@ def fp4_w4a16_linear_bf16(
191
  )
192
 
193
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
194
  def nvfp4_gemm_residual_bf16(
195
  a_packed: torch.Tensor,
196
  b_packed: torch.Tensor,
@@ -296,8 +345,10 @@ __all__ = [
296
  "fp4_w4a16_linear_bf16",
297
  "fp4_w4a4_gemv_warpsplit_bf16",
298
  "nvfp4_gemm_bf16",
 
299
  "nvfp4_gemm_bias_gelu_bf16",
300
  "nvfp4_gemm_bias_gelu_nvfp4",
 
301
  "nvfp4_gemm_residual_bf16",
302
  "nvfp4_gemm_streamk_bf16",
303
  "nvfp4_gemm_streamk_bias_bf16",
 
58
  return None
59
 
60
 
61
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_bf16"))
62
+ def _bias_fake(a, b, sfa, sfb, bias, out) -> None:
63
+ return None
64
+
65
+
66
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_residual_bf16"))
67
+ def _bias_residual_fake(a, b, sfa, sfb, bias, residual, out) -> None:
68
+ return None
69
+
70
+
71
  @torch.library.register_fake(add_op_namespace_prefix("quantize_fp4_sfa_fp16"))
72
  def _quant_fake(x: torch.Tensor, packed: torch.Tensor, sfa: torch.Tensor, is_sfb: bool = False) -> None:
73
  return None
 
201
  )
202
 
203
 
204
+ def nvfp4_gemm_bias_bf16(
205
+ a_packed: torch.Tensor,
206
+ b_packed: torch.Tensor,
207
+ sfa: torch.Tensor,
208
+ sfb: torch.Tensor,
209
+ bias: torch.Tensor,
210
+ *,
211
+ out: torch.Tensor | None = None,
212
+ ) -> torch.Tensor:
213
+ """SM110 NVFP4 GEMM with a fused per-column BF16 bias."""
214
+ if out is None:
215
+ out = torch.empty(
216
+ (a_packed.shape[0], b_packed.shape[0]),
217
+ device=a_packed.device,
218
+ dtype=torch.bfloat16,
219
+ )
220
+ ops.nvfp4_gemm_bias_bf16(a_packed, b_packed, sfa, sfb, bias, out)
221
+ return out
222
+
223
+
224
+ def nvfp4_gemm_bias_residual_bf16(
225
+ a_packed: torch.Tensor,
226
+ b_packed: torch.Tensor,
227
+ sfa: torch.Tensor,
228
+ sfb: torch.Tensor,
229
+ bias: torch.Tensor,
230
+ residual: torch.Tensor,
231
+ *,
232
+ out: torch.Tensor | None = None,
233
+ ) -> torch.Tensor:
234
+ """SM110 NVFP4 GEMM with fused BF16 bias and residual add."""
235
+ if out is None:
236
+ out = torch.empty_like(residual)
237
+ ops.nvfp4_gemm_bias_residual_bf16(
238
+ a_packed, b_packed, sfa, sfb, bias, residual, out
239
+ )
240
+ return out
241
+
242
+
243
  def nvfp4_gemm_residual_bf16(
244
  a_packed: torch.Tensor,
245
  b_packed: torch.Tensor,
 
345
  "fp4_w4a16_linear_bf16",
346
  "fp4_w4a4_gemv_warpsplit_bf16",
347
  "nvfp4_gemm_bf16",
348
+ "nvfp4_gemm_bias_bf16",
349
  "nvfp4_gemm_bias_gelu_bf16",
350
  "nvfp4_gemm_bias_gelu_nvfp4",
351
+ "nvfp4_gemm_bias_residual_bf16",
352
  "nvfp4_gemm_residual_bf16",
353
  "nvfp4_gemm_streamk_bf16",
354
  "nvfp4_gemm_streamk_bias_bf16",
build/torch213-cxx11-cu130-x86_64-linux/{_fp4_gemm_cuda_0fe564a.abi3.so → _fp4_gemm_cuda_8a66d8b.abi3.so} RENAMED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:2d9cbd6af0422c5d49dd2d98ac57a58de7834902e280e1c533af843faa35700c
3
- size 2428392
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2fe7efef250feb7e8ec881de0919f010812cc5840a2b071993b4336cdfa9717e
3
+ size 2847968
build/torch213-cxx11-cu130-x86_64-linux/_ops.py CHANGED
@@ -1,9 +1,9 @@
1
  import torch
2
- from . import _fp4_gemm_cuda_0fe564a
3
- ops = torch.ops._fp4_gemm_cuda_0fe564a
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
- return f"_fp4_gemm_cuda_0fe564a::{op_name}"
 
1
  import torch
2
+ from . import _fp4_gemm_cuda_8a66d8b
3
+ ops = torch.ops._fp4_gemm_cuda_8a66d8b
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
+ return f"_fp4_gemm_cuda_8a66d8b::{op_name}"
build/torch213-cxx11-cu130-x86_64-linux/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "fp4-gemm",
3
- "id": "_fp4_gemm_cuda_0fe564a",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
@@ -14,19 +14,19 @@
14
  "digest": {
15
  "algorithm": "sha256",
16
  "files": {
17
- "__init__.py": "pJlGLDlOKnju+m59WyZJjpynMTvkPQTF6w+z4YfCwSY=",
18
- "_fp4_gemm_cuda_0fe564a.abi3.so": "LZy9avBCLF1J3S2YrFeljeeDSQLigOHFM6+EP6o1cAw=",
19
- "_ops.py": "64wwcwE28ue+jCMSXh8vVgbUevYtei6uGZjpzaoMVrw="
20
  }
21
  },
22
  "provenance": {
23
  "kernel-builder": {
24
  "version": "0.17.0-dev0",
25
- "sha": "d720fa90fb9cd92d1bc60a9dc5c55bef2aafabb8",
26
  "dirty": false
27
  },
28
  "kernel": {
29
- "sha": "0fe564afa690cd8a731729a59f98be95ebee4b34",
30
  "dirty": false
31
  }
32
  }
 
1
  {
2
  "name": "fp4-gemm",
3
+ "id": "_fp4_gemm_cuda_8a66d8b",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
 
14
  "digest": {
15
  "algorithm": "sha256",
16
  "files": {
17
+ "__init__.py": "+Kk/VNnIWwe9nszQYvL2qtU3DpKxbqHVfklg5WtVuB4=",
18
+ "_fp4_gemm_cuda_8a66d8b.abi3.so": "L+fv7yUP636OyIHeCRnwEIEsxYQKKwcZk7QzbN+pcX4=",
19
+ "_ops.py": "mZSHUC0K9o9glzvGmRSdtPppOq4bmIaxNEjibkaWTk0="
20
  }
21
  },
22
  "provenance": {
23
  "kernel-builder": {
24
  "version": "0.17.0-dev0",
25
+ "sha": "870e825d881664e39f9287a27a74ef63ff3c545e",
26
  "dirty": false
27
  },
28
  "kernel": {
29
+ "sha": "8a66d8b7f79a6fde86aa6929db2b60ab097e2d55",
30
  "dirty": false
31
  }
32
  }
build/torch213-cxx11-cu132-x86_64-linux/__init__.py CHANGED
@@ -58,6 +58,16 @@ def _legacy_linear_fake(
58
  return None
59
 
60
 
 
 
 
 
 
 
 
 
 
 
61
  @torch.library.register_fake(add_op_namespace_prefix("quantize_fp4_sfa_fp16"))
62
  def _quant_fake(x: torch.Tensor, packed: torch.Tensor, sfa: torch.Tensor, is_sfb: bool = False) -> None:
63
  return None
@@ -191,6 +201,45 @@ def fp4_w4a16_linear_bf16(
191
  )
192
 
193
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
194
  def nvfp4_gemm_residual_bf16(
195
  a_packed: torch.Tensor,
196
  b_packed: torch.Tensor,
@@ -296,8 +345,10 @@ __all__ = [
296
  "fp4_w4a16_linear_bf16",
297
  "fp4_w4a4_gemv_warpsplit_bf16",
298
  "nvfp4_gemm_bf16",
 
299
  "nvfp4_gemm_bias_gelu_bf16",
300
  "nvfp4_gemm_bias_gelu_nvfp4",
 
301
  "nvfp4_gemm_residual_bf16",
302
  "nvfp4_gemm_streamk_bf16",
303
  "nvfp4_gemm_streamk_bias_bf16",
 
58
  return None
59
 
60
 
61
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_bf16"))
62
+ def _bias_fake(a, b, sfa, sfb, bias, out) -> None:
63
+ return None
64
+
65
+
66
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_residual_bf16"))
67
+ def _bias_residual_fake(a, b, sfa, sfb, bias, residual, out) -> None:
68
+ return None
69
+
70
+
71
  @torch.library.register_fake(add_op_namespace_prefix("quantize_fp4_sfa_fp16"))
72
  def _quant_fake(x: torch.Tensor, packed: torch.Tensor, sfa: torch.Tensor, is_sfb: bool = False) -> None:
73
  return None
 
201
  )
202
 
203
 
204
+ def nvfp4_gemm_bias_bf16(
205
+ a_packed: torch.Tensor,
206
+ b_packed: torch.Tensor,
207
+ sfa: torch.Tensor,
208
+ sfb: torch.Tensor,
209
+ bias: torch.Tensor,
210
+ *,
211
+ out: torch.Tensor | None = None,
212
+ ) -> torch.Tensor:
213
+ """SM110 NVFP4 GEMM with a fused per-column BF16 bias."""
214
+ if out is None:
215
+ out = torch.empty(
216
+ (a_packed.shape[0], b_packed.shape[0]),
217
+ device=a_packed.device,
218
+ dtype=torch.bfloat16,
219
+ )
220
+ ops.nvfp4_gemm_bias_bf16(a_packed, b_packed, sfa, sfb, bias, out)
221
+ return out
222
+
223
+
224
+ def nvfp4_gemm_bias_residual_bf16(
225
+ a_packed: torch.Tensor,
226
+ b_packed: torch.Tensor,
227
+ sfa: torch.Tensor,
228
+ sfb: torch.Tensor,
229
+ bias: torch.Tensor,
230
+ residual: torch.Tensor,
231
+ *,
232
+ out: torch.Tensor | None = None,
233
+ ) -> torch.Tensor:
234
+ """SM110 NVFP4 GEMM with fused BF16 bias and residual add."""
235
+ if out is None:
236
+ out = torch.empty_like(residual)
237
+ ops.nvfp4_gemm_bias_residual_bf16(
238
+ a_packed, b_packed, sfa, sfb, bias, residual, out
239
+ )
240
+ return out
241
+
242
+
243
  def nvfp4_gemm_residual_bf16(
244
  a_packed: torch.Tensor,
245
  b_packed: torch.Tensor,
 
345
  "fp4_w4a16_linear_bf16",
346
  "fp4_w4a4_gemv_warpsplit_bf16",
347
  "nvfp4_gemm_bf16",
348
+ "nvfp4_gemm_bias_bf16",
349
  "nvfp4_gemm_bias_gelu_bf16",
350
  "nvfp4_gemm_bias_gelu_nvfp4",
351
+ "nvfp4_gemm_bias_residual_bf16",
352
  "nvfp4_gemm_residual_bf16",
353
  "nvfp4_gemm_streamk_bf16",
354
  "nvfp4_gemm_streamk_bias_bf16",
build/torch213-cxx11-cu132-x86_64-linux/{_fp4_gemm_cuda_0fe564a.abi3.so → _fp4_gemm_cuda_8a66d8b.abi3.so} RENAMED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:b30ed92af61d4f1ac4f640a157d783e82a262857d7983f9d50aa1ebd738552e5
3
- size 2428344
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7043e89c27719838b2dfab5d8c82e8098528a792ba5d89861a5bfe45cecc022c
3
+ size 2843824
build/torch213-cxx11-cu132-x86_64-linux/_ops.py CHANGED
@@ -1,9 +1,9 @@
1
  import torch
2
- from . import _fp4_gemm_cuda_0fe564a
3
- ops = torch.ops._fp4_gemm_cuda_0fe564a
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
- return f"_fp4_gemm_cuda_0fe564a::{op_name}"
 
1
  import torch
2
+ from . import _fp4_gemm_cuda_8a66d8b
3
+ ops = torch.ops._fp4_gemm_cuda_8a66d8b
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
+ return f"_fp4_gemm_cuda_8a66d8b::{op_name}"
build/torch213-cxx11-cu132-x86_64-linux/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "fp4-gemm",
3
- "id": "_fp4_gemm_cuda_0fe564a",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
@@ -14,19 +14,19 @@
14
  "digest": {
15
  "algorithm": "sha256",
16
  "files": {
17
- "__init__.py": "pJlGLDlOKnju+m59WyZJjpynMTvkPQTF6w+z4YfCwSY=",
18
- "_fp4_gemm_cuda_0fe564a.abi3.so": "sw7ZKvYdTxrE9kChV9eD6ComKFfXmD+dUKoevXOFUuU=",
19
- "_ops.py": "64wwcwE28ue+jCMSXh8vVgbUevYtei6uGZjpzaoMVrw="
20
  }
21
  },
22
  "provenance": {
23
  "kernel-builder": {
24
  "version": "0.17.0-dev0",
25
- "sha": "d720fa90fb9cd92d1bc60a9dc5c55bef2aafabb8",
26
  "dirty": false
27
  },
28
  "kernel": {
29
- "sha": "0fe564afa690cd8a731729a59f98be95ebee4b34",
30
  "dirty": false
31
  }
32
  }
 
1
  {
2
  "name": "fp4-gemm",
3
+ "id": "_fp4_gemm_cuda_8a66d8b",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
 
14
  "digest": {
15
  "algorithm": "sha256",
16
  "files": {
17
+ "__init__.py": "+Kk/VNnIWwe9nszQYvL2qtU3DpKxbqHVfklg5WtVuB4=",
18
+ "_fp4_gemm_cuda_8a66d8b.abi3.so": "cEPonCdxmDiy36tdjILoCYUop5K6XYmGGlv+Rc7MAiw=",
19
+ "_ops.py": "mZSHUC0K9o9glzvGmRSdtPppOq4bmIaxNEjibkaWTk0="
20
  }
21
  },
22
  "provenance": {
23
  "kernel-builder": {
24
  "version": "0.17.0-dev0",
25
+ "sha": "870e825d881664e39f9287a27a74ef63ff3c545e",
26
  "dirty": false
27
  },
28
  "kernel": {
29
+ "sha": "8a66d8b7f79a6fde86aa6929db2b60ab097e2d55",
30
  "dirty": false
31
  }
32
  }