class FP8LinearLayer: """Example of integrating DeepGEMM into a training loop""" def __init__(self, in_features, out_features): self.weight = torch.randn((out_features, in_features), device='cuda', dtype=torch.bfloat16) self.weight_grad = torch.zeros_like(self.weight, dtype=torch.float) def forward(self, x): # Convert inputs to FP8 x_fp8, x_scales = cast_to_fp8_per_token(x) w_fp8, w_scales = cast_to_fp8_per_block(self.weight) # Prepare DeepGEMM inputs x_input = (x_fp8, get_col_major_tma_aligned_tensor(x_scales)) w_input = (w_fp8, w_scales) # Allocate output output = torch.empty((x.shape[0], self.weight.shape[0]), device='cuda', dtype=torch.bfloat16) # Forward pass deep_gemm.gemm_fp8_fp8_bf16_nt(x_input, w_input, output) return output def backward(self, x, grad_output): # Convert to FP8 x_fp8, x_scales = cast_to_fp8_per_token(x) grad_fp8, grad_scales = cast_to_fp8_per_token(grad_output) # Prepare inputs x_input = (x_fp8, get_col_major_tma_aligned_tensor(x_scales)) grad_input = (grad_fp8, get_col_major_tma_aligned_tensor(grad_scales)) # Compute weight gradients: grad_output^T @ x deep_gemm.wgrad_gemm_fp8_fp8_fp32_nt(grad_input, x_input, self.weight_grad) # Demo usage layer = FP8LinearLayer(512, 256) x = torch.randn((128, 512), device='cuda', dtype=torch.bfloat16) # Forward pass y = layer.forward(x) print(f"Forward output shape: {y.shape}") # Backward pass grad_y = torch.randn_like(y) layer.backward(x, grad_y) print(f"Weight grad shape: {layer.weight_grad.shape}") print("✓ Training loop integration demonstrated")