# DeepGEMM requires specific tensor layouts from deep_gemm import get_col_major_tma_aligned_tensor # LHS scales must be transposed and TMA-aligned lhs_scales_aligned = get_col_major_tma_aligned_tensor(lhs_scales) # RHS scales must be contiguous assert rhs_scales.is_contiguous() # Package the inputs lhs_input = (lhs_fp8, lhs_scales_aligned) rhs_input = (rhs_fp8, rhs_scales) print("✓ Tensors prepared for DeepGEMM") print(f"LHS scales alignment: {lhs_scales_aligned.stride()}")