When inheriting from compressai.models.CompressionModel, you can simplify the training loop by separating the compression network parameters from the entropy bottleneck (quantiles) parameters. This allows you to use two different optimizers.
To identify the parameter groups:
- Compression parameters: All parameters except those ending in
.quantiles. - Auxiliary parameters: Parameters ending in
.quantiles.
Example training loop structure:
- Zero both optimizers.
- Forward pass to get $\hat{x}$ and
y_likelihoods. - Compute and backpropagate the rate-distortion loss using the main optimizer.
- Compute and backpropagate the auxiliary loss using the auxiliary optimizer.
from compressai.models import CompressionModel
from compressai.models.utils import conv, deconv
class Network(CompressionModel):
def __init__(self, N=128):
super().__init__()
self.encode = nn.Sequential(
conv(3, N),
GDN(N),
conv(N, N),
GDN(N),
conv(N, N),
)
self.decode = nn.Sequential(
deconv(N, N),
GDN(N, inverse=True),
deconv(N, N),
GDN(N, inverse=True),
deconv(N, 3),
)
def forward(self, x):
y = self.encode(x)
y_hat, y_likelihoods = self.entropy_bottleneck(y)
x_hat = self.decode(y_hat)
return x_hat, y_likelihoods
# Optimizer setup
parameters = set(p for n, p in net.named_parameters() if not n.endswith(".quantiles"))
aux_parameters = set(p for n, p in net.named_parameters() if n.endswith(".quantiles"))
optimizer = optim.Adam(parameters, lr=1e-4)
aux_optimizer = optim.Adam(aux_parameters, lr=1e-3)
# Training loop
x = torch.rand(1, 3, 64, 64)
for i in range(10):
optimizer.zero_grad()
aux_optimizer.zero_grad()
x_hat, y_likelihoods = net(x)
# ... compute loss ...
loss.backward()
optimizer.step()
aux_loss = net.aux_loss()
aux_loss.backward()
aux_optimizer.step()