如何在PyTorch中约束自编码器卷积核仅取-1、0或1值?
Absolutely! You can absolutely constrain your convolutional kernel weights to only take values of -1, 0, or 1 in PyTorch—and this is a great idea given your input images use discrete {-1,0,1} pixel values. Aligning your model's weights with this discrete distribution could help pull it out of those float-valued local minima by forcing it to learn patterns that match your input's structure better.
Here are two straightforward ways to implement this:
1. 训练循环中手动约束权重
You can define a helper function that projects your convolutional weights to {-1,0,1}, then call it right after each optimizer step to enforce the constraint:
import torch import torch.nn as nn def constrain_conv_weights(model): # Iterate over all layers in the model for m in model.modules(): # Target only convolutional and transposed convolutional layers if isinstance(m, (nn.Conv2d, nn.ConvTranspose2d)): with torch.no_grad(): # No need to track gradients for this operation # Project weights to {-1, 0, 1}: # Values > 0.5 become 1, < -0.5 become -1, everything else becomes 0 m.weight.data = torch.where( m.weight.data > 0.5, torch.tensor(1.0, device=m.weight.device), torch.where( m.weight.data < -0.5, torch.tensor(-1.0, device=m.weight.device), torch.tensor(0.0, device=m.weight.device) ) )
Then in your training loop, add this right after updating the weights:
# Inside your training loop optimizer.zero_grad() recon = model(inputs) loss = loss_fn(recon, inputs) loss.backward() optimizer.step() # Enforce weight constraint after each update constrain_conv_weights(model)
2. 自定义离散卷积层
For a more encapsulated approach, you can create a custom convolutional layer that automatically applies the weight constraint before every forward pass:
class DiscreteConv2d(nn.Conv2d): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) def forward(self, x): # Constrain weights to {-1,0,1} before computing the convolution with torch.no_grad(): self.weight.data = torch.where( self.weight.data > 0.5, torch.tensor(1.0, device=self.weight.device), torch.where( self.weight.data < -0.5, torch.tensor(-1.0, device=self.weight.device), torch.tensor(0.0, device=self.weight.device) ) ) return super().forward(x) # Do the same for transposed convolutions if you use them in the decoder class DiscreteConvTranspose2d(nn.ConvTranspose2d): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) def forward(self, x): with torch.no_grad(): self.weight.data = torch.where( self.weight.data > 0.5, torch.tensor(1.0, device=self.weight.device), torch.where( self.weight.data < -0.5, torch.tensor(-1.0, device=self.weight.device), torch.tensor(0.0, device=self.weight.device) ) ) return super().forward(x)
Then when building your autoencoder, use DiscreteConv2d and DiscreteConvTranspose2d instead of the standard PyTorch layers—no need to add extra code in your training loop!
关键注意事项
- Adjust your learning rate: Since we're hard-constraining weights to discrete values, a smaller learning rate (e.g., 1e-4 instead of 1e-3) will help prevent the optimizer from bouncing weights between values too aggressively.
- Consider pre-training first: You might want to train the model without constraints for a few epochs to let it learn basic features, then switch on the weight constraints for fine-tuning. This can help speed up convergence.
- Pair with discrete output handling: Since your inputs are {-1,0,1}, you could also modify your decoder's final layer to output discrete values (e.g., using
torch.signafter a linear layer, or a custom activation that maps to {-1,0,1}). This would align your output distribution even better with your inputs.
内容的提问来源于stack exchange,提问作者Priva

