PyTorch torch.matmul function


Pytorch torch 参考手册Pytorch torch reference manual

torch.matmulIt is a function in PyTorch used to perform matrix multiplication. It supports inputs of different dimensions and can handle tensor multiplication of one-dimensional, two-dimensional, and higher-dimensional tensors.

This is one of the most commonly used operations in deep learning. The forward propagation of a neural network is essentially a series of matrix multiplications.

Function definition

torch.matmul(input, other, out=None)

Parameters:

  • input(Tensor): The first input tensor.
  • other(Tensor): The second input tensor.
  • out(Tensor, optional): The output tensor.

Return value:

  • torch.Tensor: Returns the result of matrix multiplication.

Usage examples

Example 1: Two-dimensional matrix multiplication

Example

import torch

# Create two 2D matrices
a = torch.randn(3, 4)  # 3x4 matrix
b = torch.randn(4, 5)  # 4x5 matrix

# Matrix multiplication
c = torch.matmul(a, b)

print("Shape of a:", a.shape)
print("Shape of b:", b.shape)
print("Shape of c:", c.shape)

The output is:

a 的形状: torch.Size([3, 4])
b 的形状: torch.Size([4, 5])
c 的形状: torch.Size([3, 5])

Example 2: Multiplying a vector and a matrix

Example

import torch

# Create a vector and a matrix
vector = torch.randn(4)      # 4-dimensional vector
matrix = torch.randn(4, 5)  # 4x5 matrix

# Multiply vector and matrix
result = torch.matmul(vector, matrix)

print("Shape of the vector:", vector.shape)
print("Shape of the matrix:", matrix.shape)
print("Shape of the result:", result.shape)
print(result)

The output is:

向量的形状: torch.Size([4])
矩阵的形状: torch.Size([4, 5])
结果的形状: torch.Size([5])
tensor([-0.7837,  0.3684, -0.6542, -0.4594,  1.5328])
</p>

<h3>示例 3: 批量矩阵乘法</h3>

<div class="example">
<h2 class="example">实例</h2>
<div class="example_code">
<span style="color: Green;font-weight:bold;">import</span> torch<br />
<br />
<span style="color: #a50"># 创建批量矩阵</span><br />
batch_a <span style="color: Gray;">=</span> torch.<span style="color: #05a;">randn</span><span style="color: Olive;">&#40;</span><span style="color: Maroon;">10</span><span style="color: Gray;">,</span> <span style="color: Maroon;">3</span><span style="color: Gray;">,</span> <span style="color: Maroon;">4</span><span style="color: Olive;">&#41;</span> &nbsp;<span style="color: #a50"># 10 个 3x4 矩阵</span><br />
batch_b <span style="color: Gray;">=</span> torch.<span style="color: #05a;">randn</span><span style="color: Olive;">&#40;</span><span style="color: Maroon;">10</span><span style="color: Gray;">,</span> <span style="color: Maroon;">4</span><span style="color: Gray;">,</span> <span style="color: Maroon;">5</span><span style="color: Olive;">&#41;</span> &nbsp;<span style="color: #a50"># 10 个 4x5 矩阵</span><br />
<br />
<span style="color: #a50"># 批量矩阵乘法</span><br />
batch_c <span style="color: Gray;">=</span> torch.<span style="color: #05a;">matmul</span><span style="color: Olive;">&#40;</span>batch_a<span style="color: Gray;">,</span> batch_b<span style="color: Olive;">&#41;</span><br />
<br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">&#40;</span><span style="color: #a11;">&quot;批量 a 的形状:&quot;</span><span style="color: Gray;">,</span> batch_a.<span style="color: #05a;">shape</span><span style="color: Olive;">&#41;</span><br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">&#40;</span><span style="color: #a11;">&quot;批量 b 的形状:&quot;</span><span style="color: Gray;">,</span> batch_b.<span style="color: #05a;">shape</span><span style="color: Olive;">&#41;</span><br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">&#40;</span><span style="color: #a11;">&quot;批量结果 c 的形状:&quot;</span><span style="color: Gray;">,</span> batch_c.<span style="color: #05a;">shape</span><span style="color: Olive;">&#41;</span><br />
</div>
</div>

<p>输出结果为:</p>
<pre>
批量 a 的形状: torch.Size([10, 3, 4])
批量 b 的形状: torch.Size([10, 4, 5])
批量结果 c 的形状: torch.Size([10, 3, 5])

Batched matrix multiplication is a common operation in deep learning, for example in the attention mechanism of Transformers.

Example 4: Matrix multiplication in neural networks

Example

import torch

# Simulate a neural network layer: input xW + b
x = torch.randn(32, 128)   # Batch size 32, feature dimension 128
W = torch.randn(128, 256)  # Weight matrix
b = torch.randn(256)       # Bias vector

# Linear transformation: y = x @ W^T + b (in PyTorch, W is usually transposed)
# Here we demonstrate x @ W
y = torch.matmul(x, W) + b

print("Shape of input x:", x.shape)
print("Shape of weight W:", W.shape)
print("Shape of output y:", y.shape)

The output is:

输入 x 形状: torch.Size([32, 128])
权重 W 形状: torch.Size([128, 256])
输出 y 形状: torch.Size([32, 256])

This example simulates the forward propagation process of a fully connected layer in a neural network.


Notes

  • torch.matmulBroadcasting is supported, but the last two dimensions of both inputs must satisfy the dimension requirements of matrix multiplication.
  • For two-dimensional matrix multiplication, you can also usetorch.mm(), buttorch.matmulis more general.
  • Note the difference betweentorch.matmul(matrix multiplication) andtorch.mul(element-wise multiplication).

Pytorch torch 参考手册Pytorch torch reference manual

Other extensions