PyTorch torch.nn.Bilinear Function
PyTorch torch.nn Reference Manual
torch.nn.BilinearIt is a bilinear layer in PyTorch.
It performs a bilinear transformation on two inputs, commonly used for feature fusion.
Function Definition
torch.nn.Bilinear(in1_features, out_features, out_features, bias=True)
Formula
y = x1 * W * x2 + b
Usage Example
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
# Bilinear layer: two 100-dimensional inputs -> 50-dimensional output
bilinear = nn.Bilinear(100, 100, 50)
# Two inputs
x1 = torch.randn(4, 100)
x2 = torch.randn(4, 100)
output = bilinear(x1, x2)
print("Input 1:", x1.shape)
print("Input 2:", x2.shape)
print("Output:", output.shape)
import torch.nn as nn
# Bilinear layer: two 100-dimensional inputs -> 50-dimensional output
bilinear = nn.Bilinear(100, 100, 50)
# Two inputs
x1 = torch.randn(4, 100)
x2 = torch.randn(4, 100)
output = bilinear(x1, x2)
print("Input 1:", x1.shape)
print("Input 2:", x2.shape)
print("Output:", output.shape)
Example 2: Feature Fusion
Example
import torch
import torch.nn as nn
# Bilinear feature fusion (similar to Skip-gram)
class FusionNet(nn.Module):
def __init__(self, dim=128):
super(FusionNet, self).__init__()
self.bilinear = nn.Bilinear(dim, dim, dim)
def forward(self, feat1, feat2):
return self.bilinear(feat1, feat2)
model = FusionNet()
f1 = torch.randn(4, 128)
f2 = torch.randn(4, 128)
output = model(f1, f2)
print("Fusion output:", output.shape)
import torch.nn as nn
# Bilinear feature fusion (similar to Skip-gram)
class FusionNet(nn.Module):
def __init__(self, dim=128):
super(FusionNet, self).__init__()
self.bilinear = nn.Bilinear(dim, dim, dim)
def forward(self, feat1, feat2):
return self.bilinear(feat1, feat2)
model = FusionNet()
f1 = torch.randn(4, 128)
f2 = torch.randn(4, 128)
output = model(f1, f2)
print("Fusion output:", output.shape)
Example 3: Parameter Count
Example
import torch
import torch.nn as nn
# Bilinear layer parameters
bilinear = nn.Bilinear(100, 100, 50)
print("Number of parameters:", sum(p.numel() for p in bilinear.parameters()))
print("Weight shape:", bilinear.weight.shape) # (50, 100, 100)
import torch.nn as nn
# Bilinear layer parameters
bilinear = nn.Bilinear(100, 100, 50)
print("Number of parameters:", sum(p.numel() for p in bilinear.parameters()))
print("Weight shape:", bilinear.weight.shape) # (50, 100, 100)
Use Cases
- Feature fusion: Multimodal
- Attention mechanism
- Interaction modeling
Note: The bilinear layer has a large number of parameters; use with caution.
Other Extensions