CLIP-Search-Edit / cnn1.py
randomoneguy's picture
new-cnn (#1)
5a14c00 verified
Raw
History Blame Contribute Delete
7.88 kB
import torch
import torch.nn as nn
class FiLMLayer(nn.Module):
"""
Feature-wise linear modulation module that conditions convolutional activations
on an external style embedding (e.g., a CLIP text embedding).
"""
def __init__(self, num_channels: int, cond_dim: int, hidden_dim: int = 256):
super().__init__()
self.net = nn.Sequential(
nn.LayerNorm(cond_dim),
nn.Linear(cond_dim, hidden_dim),
nn.GELU(),
nn.Linear(hidden_dim, num_channels * 2),
)
def forward(self, x: torch.Tensor, cond: torch.Tensor) -> torch.Tensor:
gamma, beta = self.net(cond).chunk(2, dim=1)
gamma = gamma.unsqueeze(-1).unsqueeze(-1)
beta = beta.unsqueeze(-1).unsqueeze(-1)
return x * (1 + gamma) + beta
class DoubleConv(nn.Module):
"""Two consecutive conv-batchnorm-gelu blocks with optional FiLM conditioning."""
def __init__(
self,
in_channels: int,
out_channels: int,
cond_dim: int | None = None,
film_hidden_dim: int = 256,
):
super().__init__()
self.conv = nn.Sequential(
nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),
nn.BatchNorm2d(out_channels),
nn.GELU(),
nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),
nn.BatchNorm2d(out_channels),
nn.GELU(),
)
self.film = (
FiLMLayer(out_channels, cond_dim, hidden_dim=film_hidden_dim)
if cond_dim is not None
else None
)
def forward(self, x: torch.Tensor, cond: torch.Tensor | None = None) -> torch.Tensor:
x = self.conv(x)
if self.film is not None:
if cond is None:
raise ValueError("Style embedding is required for FiLM conditioning.")
x = self.film(x, cond)
return x
class DownBlock(nn.Module):
"""Down-sampling block used in the encoder path."""
def __init__(
self,
in_channels: int,
out_channels: int,
cond_dim: int | None = None,
film_hidden_dim: int = 256,
):
super().__init__()
self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
self.conv = DoubleConv(
in_channels, out_channels, cond_dim, film_hidden_dim=film_hidden_dim
)
def forward(self, x: torch.Tensor, cond: torch.Tensor | None = None) -> torch.Tensor:
x = self.pool(x)
return self.conv(x, cond)
class UpBlock(nn.Module):
"""Up-sampling block with skip connections from the encoder path."""
def __init__(
self,
in_channels: int,
skip_channels: int,
cond_dim: int | None = None,
bilinear: bool = True,
film_hidden_dim: int = 256,
):
super().__init__()
if bilinear:
self.up = nn.Sequential(
nn.Upsample(scale_factor=2, mode="bilinear", align_corners=True),
nn.Conv2d(in_channels, in_channels // 2, kernel_size=1),
)
else:
self.up = nn.ConvTranspose2d(
in_channels, in_channels // 2, kernel_size=2, stride=2
)
self.conv = DoubleConv(
in_channels // 2 + skip_channels,
skip_channels,
cond_dim,
film_hidden_dim=film_hidden_dim,
)
def forward(
self, x: torch.Tensor, skip: torch.Tensor, cond: torch.Tensor | None = None
) -> torch.Tensor:
x = self.up(x)
diff_y = skip.size(2) - x.size(2)
diff_x = skip.size(3) - x.size(3)
if diff_y != 0 or diff_x != 0:
x = nn.functional.pad(
x,
[
diff_x // 2,
diff_x - diff_x // 2,
diff_y // 2,
diff_y - diff_y // 2,
],
)
x = torch.cat([skip, x], dim=1)
return self.conv(x, cond)
class OutConv(nn.Module):
"""Final projection into the RGB space."""
def __init__(self, in_channels: int, out_channels: int):
super().__init__()
self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.conv(x)
class UNet(nn.Module):
"""
Lightweight encoder-decoder network for CLIP-guided, text-prompted style transfer.
The network takes a content image and optionally a CLIP text embedding that
modulates intermediate activations through FiLM layers so that the decoded
image aligns with the target style semantics in CLIP space.
"""
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
base_channels: int = 16,
num_layers: int = 4,
text_dim: int = None,
bilinear: bool = True,
film_hidden_dim: int = 256,
):
super().__init__()
if num_layers < 2:
raise ValueError("num_layers must be >= 2")
self.cond_dim = text_dim
self.style_mapper = (
nn.Sequential(
nn.LayerNorm(text_dim),
nn.Linear(text_dim, text_dim),
nn.GELU(),
nn.Linear(text_dim, text_dim),
)
if text_dim is not None
else None
)
channels = [base_channels * (2**i) for i in range(num_layers)]
self.inc = DoubleConv(
in_channels,
channels[0],
self.cond_dim,
film_hidden_dim=film_hidden_dim,
)
self.downs = nn.ModuleList()
for idx in range(num_layers - 1):
self.downs.append(
DownBlock(
channels[idx],
channels[idx + 1],
self.cond_dim,
film_hidden_dim=film_hidden_dim,
)
)
self.bottleneck = DoubleConv(
channels[-1],
channels[-1] * 2,
self.cond_dim,
film_hidden_dim=film_hidden_dim,
)
self.ups = nn.ModuleList()
prev_channels = channels[-1] * 2
for skip_ch in reversed(channels):
self.ups.append(
UpBlock(
prev_channels,
skip_ch,
self.cond_dim,
bilinear=bilinear,
film_hidden_dim=film_hidden_dim,
)
)
prev_channels = skip_ch
self.outc = OutConv(channels[0], out_channels)
self.activation = nn.Tanh()
def _prepare_condition(self, text_embedding: torch.Tensor | None) -> torch.Tensor | None:
if self.cond_dim is None:
return None
if text_embedding is None:
raise ValueError(
"text_embedding must be provided when the model is configured for conditioning."
)
if text_embedding.dim() != 2 or text_embedding.size(1) != self.cond_dim:
raise ValueError(
f"text_embedding must have shape [batch, {self.cond_dim}] but got {text_embedding.shape}."
)
return self.style_mapper(text_embedding) if self.style_mapper else text_embedding
def forward(
self, x: torch.Tensor, text_embedding: torch.Tensor | None = None
) -> torch.Tensor:
cond = self._prepare_condition(text_embedding)
skip_connections = []
x = self.inc(x, cond)
skip_connections.append(x)
for down in self.downs:
x = down(x, cond)
skip_connections.append(x)
x = self.bottleneck(x, cond)
for up, skip in zip(self.ups, reversed(skip_connections)):
x = up(x, skip, cond)
x = self.outc(x)
return self.activation(x)