# Control SM utilization for better efficiency original_sms = deep_gemm.get_num_sms() print(f"Default SMs: {original_sms}") # Use fewer SMs for smaller problems to save power deep_gemm.set_num_sms(original_sms // 2) print(f"Reduced SMs: {deep_gemm.get_num_sms()}") # Run a smaller GEMM small_lhs = torch.randn((64, 256), device='cuda', dtype=torch.bfloat16) small_rhs = torch.randn((128, 256), device='cuda', dtype=torch.bfloat16) small_out = torch.empty((64, 128), device='cuda', dtype=torch.bfloat16) small_lhs_fp8, small_lhs_scales = cast_to_fp8_per_token(small_lhs) small_rhs_fp8, small_rhs_scales = cast_to_fp8_per_block(small_rhs) deep_gemm.gemm_fp8_fp8_bf16_nt( (small_lhs_fp8, get_col_major_tma_aligned_tensor(small_lhs_scales)), (small_rhs_fp8, small_rhs_scales), small_out ) # Restore original setting deep_gemm.set_num_sms(original_sms) print("✓ SM configuration demonstrated")