import torch import deep_gemm # Create simple input matrices m, n, k = 128, 256, 512 lhs = torch.randn((m, k), device='cuda', dtype=torch.bfloat16) rhs = torch.randn((n, k), device='cuda', dtype=torch.bfloat16) output = torch.empty((m, n), device='cuda', dtype=torch.bfloat16) print(f"LHS shape: {lhs.shape}") # [128, 512] print(f"RHS shape: {rhs.shape}") # [256, 512] print(f"Output shape: {output.shape}") # [128, 256]