(mat_a, mat_b, scales_a, scales_b, alpha, bias=None)
| 2 | |
| 3 | |
| 4 | def _fp8_f16_accum_meta(mat_a, mat_b, scales_a, scales_b, out_dtype, bias=None): |
| 5 | del scales_a, scales_b, bias |
| 6 | return torch.empty((mat_a.shape[0], mat_b.shape[1]), dtype=out_dtype, device=mat_a.device) |
| 7 | |
| 8 | |
| 9 | FP8_F16_ACCUM_MM_AVAILABLE = hasattr(torch.ops.lightx2v_kernel, "cutlass_scaled_fp8_mm_f16_accum_sm120") |
| 10 | if FP8_F16_ACCUM_MM_AVAILABLE: |
| 11 | _fp8_f16_accum_op = torch.ops.lightx2v_kernel.cutlass_scaled_fp8_mm_f16_accum_sm120.default |