import deep_gemm # 普通 FP8 GEMM lhs = (torch.randn(4096, 7168, dtype=torch.float8_e4m3fn), torch.randn(4096, 56, dtype=torch.float32)) # [M, K], [M, K//128] rhs = (torch.randn(2112, 7168, dtype=torch.float8_e4m3fn), torch.randn(17, 56, dtype=torch.float32)) # [N, K], [N//128, K//128] out = torch.empty(4096, 2112, dtype=torch.bfloat16, device="cuda") deep_gemm.gemm_fp8_fp8_bf16_nt(lhs, rhs, out) # MoE 分组 GEMM(连续布局) m_indices = torch.randint(0, 4, (8192,), device="cuda") # 分组索引 deep_gemm.m_grouped_gemm_fp8_fp8_bf16_nt_contiguous(lhs, rhs, out, m_indices)