OpenDelight-Demo / decoder /sapiens_decoder.py
suvadityamuk's picture
suvadityamuk HF Staff
Add OpenDelight ZeroGPU demo
2327452 verified
Raw
History Blame Contribute Delete
6.69 kB
import torch
import torch.nn as nn
from torch.utils.checkpoint import checkpoint
def count_parameters(model: nn.Module, detailed: bool = True) -> int:
"""Count parameters in a PyTorch model.
Args:
model: PyTorch model to count parameters for
detailed: whether to print per-layer details
Returns:
Total number of parameters in the model
"""
from prettytable import PrettyTable
total_params = 0
if detailed:
table = PrettyTable(["Layer Name", "Param Count", "Trainable"])
for name, parameter in model.named_parameters():
if "img_encoder" in name:
continue
if "vol_decoder" in name:
continue
param_count = parameter.numel()
total_params += param_count
if detailed:
table.add_row([name, param_count, parameter.requires_grad])
if detailed:
print(table)
print(f"\nTotal parameters: {total_params:,}")
print(f"Trainable parameters: {sum(p.numel() for p in model.parameters() if p.requires_grad):,}")
return total_params
class SapiensDecoder(nn.Module):
def __init__(self,
in_channels,
out_channels,
deconv_out_channels,
deconv_kernel_sizes,
conv_out_channels,
conv_kernel_sizes,
interpolate_mode='bilinear',
grad_checkpointing=False,
**kwargs):
super().__init__()
self.interpolate_mode = interpolate_mode
self.grad_checkpointing = grad_checkpointing
if deconv_out_channels:
if deconv_kernel_sizes is None or len(deconv_out_channels) != len(
deconv_kernel_sizes):
raise ValueError(
'"deconv_out_channels" and "deconv_kernel_sizes" should '
'be integer sequences with the same length. Got '
f'mismatched lengths {deconv_out_channels} and '
f'{deconv_kernel_sizes}')
self.deconv_layers = self._make_deconv_layers(
in_channels=in_channels,
layer_out_channels=deconv_out_channels,
layer_kernel_sizes=deconv_kernel_sizes,
)
in_channels = deconv_out_channels[-1]
else:
self.deconv_layers = nn.Identity()
if conv_out_channels:
if conv_kernel_sizes is None or len(conv_out_channels) != len(
conv_kernel_sizes):
raise ValueError(
'"conv_out_channels" and "conv_kernel_sizes" should '
'be integer sequences with the same length. Got '
f'mismatched lengths {conv_out_channels} and '
f'{conv_kernel_sizes}')
self.conv_layers = self._make_conv_layers(
in_channels=in_channels,
layer_out_channels=conv_out_channels,
layer_kernel_sizes=conv_kernel_sizes)
in_channels = conv_out_channels[-1]
else:
self.conv_layers = nn.Identity()
self.cls_conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)
def _make_conv_layers(self, in_channels,
layer_out_channels,
layer_kernel_sizes) -> nn.Module:
"""Create convolutional layers by given parameters."""
layers = []
for out_channels, kernel_size in zip(layer_out_channels,
layer_kernel_sizes):
padding = (kernel_size - 1) // 2
cfg = dict(
in_channels=in_channels,
out_channels=out_channels,
kernel_size=kernel_size,
stride=1,
padding=padding)
layers.append(
nn.Conv2d(**cfg)
)
layers.append(nn.InstanceNorm2d(out_channels))
layers.append(nn.SiLU(inplace=True))
in_channels = out_channels
return nn.Sequential(*layers)
def _make_deconv_layers(self, in_channels: int,
layer_out_channels,
layer_kernel_sizes) -> nn.Module:
"""Create deconvolutional layers by given parameters."""
layers = []
for out_channels, kernel_size in zip(layer_out_channels,
layer_kernel_sizes):
if kernel_size == 4:
padding = 1
output_padding = 0
elif kernel_size == 3:
padding = 1
output_padding = 1
elif kernel_size == 2:
padding = 0
output_padding = 0
else:
raise ValueError(f'Unsupported kernel size {kernel_size} for'
'deconvlutional layers in '
f'{self.__class__.__name__}')
cfg = dict(
in_channels=in_channels,
out_channels=out_channels,
kernel_size=kernel_size,
stride=2,
padding=padding,
output_padding=output_padding,
bias=False)
layers.append(
nn.ConvTranspose2d(**cfg)
)
layers.append(nn.InstanceNorm2d(out_channels))
layers.append(nn.SiLU(inplace=True))
in_channels = out_channels
return nn.Sequential(*layers)
def forward(self, inputs):
if self.grad_checkpointing and self.training:
x = checkpoint(self.deconv_layers, inputs, use_reentrant=False) ## B x 768 x 512 x 384
x = checkpoint(self.conv_layers, x, use_reentrant=False) ## B x 768 x 512 x 384
else:
# inputs = self._transform_inputs(inputs) ## B x 1536 x 64 x 48
x = self.deconv_layers(inputs) ## B x 768 x 512 x 384
x = self.conv_layers(x) ## B x 768 x 512 x 384
out = self.cls_conv(x) ## B x 1 x 512 x 384
return out
if __name__ == "__main__":
device = "cuda"
network = SapiensDecoder(
in_channels=512,
out_channels=3,
deconv_out_channels=(512, 512, 512, 512), ## this will 2x at each step. so total is 8x
deconv_kernel_sizes=(4, 4, 4, 4),
conv_out_channels=(512, 512, 512, 512),
conv_kernel_sizes=(1, 1, 1, 1),
).to(device)
feat = torch.randn(2, 512, 32, 32).to(device)
with torch.no_grad():
output = network(feat)
print(output.shape) ## B x 1 x 256 x 192
count_parameters(network, detailed=True)