# For training: compute weight gradients def setup_weight_gradient(): m_grad, k_grad, n_grad = 256, 1024, 512 # Activations (forward pass) activations = torch.randn((m_grad, k_grad), device='cuda', dtype=torch.bfloat16) # Gradient w.r.t. output (from backprop) grad_output = torch.randn((m_grad, n_grad), device='cuda', dtype=torch.bfloat16) # Weight gradient accumulator (typically has residual) weight_grad = torch.randn((n_grad, k_grad), device='cuda', dtype=torch.float) * 0.1 return activations, grad_output, weight_grad activations, grad_output, weight_grad = setup_weight_gradient() # Convert to FP8 act_fp8, act_scales = cast_to_fp8_per_token(activations) grad_fp8, grad_scales = cast_to_fp8_per_token(grad_output) # Prepare inputs (both need transposed scales) act_input = (act_fp8, get_col_major_tma_aligned_tensor(act_scales)) grad_input = (grad_fp8, get_col_major_tma_aligned_tensor(grad_scales)) print(f"Weight gradient shape: {weight_grad.shape}") print(f"Accumulator dtype: {weight_grad.dtype}") # FP32 for precision