TensorFlow 1.4中单个矩阵与批量矩阵相乘的高效实现方法
Great question! The core solution here is using broadcasting (supported in all major ML/numerical frameworks like PyTorch, NumPy, TensorFlow) — this lets you perform the batch multiplication without explicitly duplicating matrix A in memory. Instead of creating a full (batch_size, dim1, dim2) tiled copy of A (which eats up unnecessary RAM), broadcasting reuses the original A data during computation.
How It Works
Matrix multiplication requires inner dimensions to match:
- Your 2D matrix
Ahas shape(dim1, dim2) - Your batch of matrices has shape
(batch_size, dim2, dim3)
To align these for batch multiplication, you just need to add a singleton (size-1) batch dimension to A, turning it into (1, dim1, dim2). Frameworks will automatically broadcast this singleton dimension to match the batch_size of your other matrices during multiplication—no actual data duplication happens behind the scenes.
Example Code (PyTorch)
import torch # Define sample dimensions dim1, dim2, dim3, batch_size = 5, 10, 3, 8 # Initialize your matrices A = torch.randn(dim1, dim2) # Shape: (5, 10) batch_matrices = torch.randn(batch_size, dim2, dim3) # Shape: (8, 10, 3) # Add a singleton batch dimension to A (no memory copy!) A_expanded = A.unsqueeze(0) # Equivalent to A[None, :, :], shape: (1, 5, 10) # Perform batch matrix multiplication result = torch.bmm(A_expanded, batch_matrices) # Result shape: (8, 5, 3) — each element is A multiplied by the corresponding batch matrix
Example Code (NumPy)
import numpy as np # Same sample dimensions dim1, dim2, dim3, batch_size = 5, 10, 3, 8 A = np.random.randn(dim1, dim2) # Shape: (5, 10) batch_matrices = np.random.randn(batch_size, dim2, dim3) # Shape: (8, 10, 3) # Add singleton batch dimension A_expanded = A[np.newaxis, :, :] # Shape: (1, 5, 10) # Use np.matmul (supports broadcasting for batch operations) result = np.matmul(A_expanded, batch_matrices) # Result shape: (8, 5, 3)
Key Notes
- No memory bloat: Unlike using
repeatortileto explicitly create a tiled version of A, broadcasting only uses the originaldim1*dim2memory for A. - Framework compatibility: This approach works across all major frameworks—just adjust the dimension-expansion syntax (e.g.,
tf.expand_dims(A, 0)for TensorFlow) and use the appropriate batch multiplication function (or let the framework handle it via standard matmul).
内容的提问来源于stack exchange,提问作者antande

