File size: 7,813 Bytes
4811c23 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 | import torch, types, copy
from torch import nn
import torch.nn.functional as F
from diffusers.models.unets.unet_2d_blocks import CrossAttnDownBlock2D, \
CrossAttnUpBlock2D, \
DownBlock2D, \
UpBlock2D, \
UNetMidBlock2DCrossAttn
from diffusers.models.resnet import ResnetBlock2D
from diffusers.models.transformers.transformer_2d import Transformer2DModel
from diffusers.models.attention import BasicTransformerBlock
from diffusers.models.downsampling import Downsample2D
from diffusers.models.upsampling import Upsample2D
from forward import MyUNet2DConditionModel_SD_forward, \
MyCrossAttnDownBlock2D_SD_forward, \
MyDownBlock2D_SD_forward, \
MyUNetMidBlock2DCrossAttn_SD_forward, \
MyCrossAttnUpBlock2D_SD_forward, \
MyUpBlock2D_SD_forward, \
MyResnetBlock2D_SD_forward, \
MyTransformer2DModel_SD_forward
def find_parent(model, module_name):
components = module_name.split(".")
parent = model
for comp in components[:-1]:
parent = getattr(parent, comp)
return parent, components[-1]
def halve_channels(model):
for name, module in model.named_modules():
if hasattr(module, "pruned"):
continue
if isinstance(module, nn.Conv2d):
in_channels = int(module.in_channels * 0.75)
out_channels = int(module.out_channels * 0.75)
new_conv = nn.Conv2d(in_channels=in_channels,
out_channels=out_channels,
kernel_size=module.kernel_size,
stride=module.stride,
padding=module.padding,
dilation=module.dilation,
groups=module.groups,
bias=module.bias is not None)
with torch.no_grad():
new_conv.weight.copy_(module.weight[:out_channels, :in_channels])
if module.bias is not None:
new_conv.bias.copy_(module.bias[:out_channels])
parent, last_name = find_parent(model, name)
setattr(parent, last_name, new_conv)
new_conv.pruned = True
elif isinstance(module, nn.Linear):
in_features = int(module.in_features * 0.75)
out_features = int(module.out_features * 0.75)
new_linear = nn.Linear(in_features=in_features,
out_features=out_features,
bias=module.bias is not None)
with torch.no_grad():
new_linear.weight.copy_(module.weight[:out_features, :in_features])
if module.bias is not None:
new_linear.bias.copy_(module.bias[:out_features])
parent, last_name = find_parent(model, name)
setattr(parent, last_name, new_linear)
new_linear.pruned = True
elif isinstance(module, nn.GroupNorm):
num_channels = int(module.num_channels * 0.75)
for num_groups in [32, 24, 16, 12, 8, 6, 4, 2, 1]:
if num_channels % num_groups == 0:
break
new_gn = nn.GroupNorm(num_groups=num_groups,
num_channels=num_channels,
eps=module.eps,
affine=module.affine)
with torch.no_grad():
new_gn.weight.copy_(module.weight[:num_channels])
new_gn.bias.copy_(module.bias[:num_channels])
parent, last_name = find_parent(model, name)
setattr(parent, last_name, new_gn)
new_gn.pruned = True
elif isinstance(module, nn.LayerNorm):
normalized_shape = int(module.normalized_shape[0] * 0.75)
new_ln = nn.LayerNorm(normalized_shape,
eps=module.eps,
elementwise_affine=module.elementwise_affine)
with torch.no_grad():
new_ln.weight.copy_(module.weight[:normalized_shape])
new_ln.bias.copy_(module.bias[:normalized_shape])
parent, last_name = find_parent(model, name)
setattr(parent, last_name, new_ln)
new_ln.pruned = True
elif isinstance(module, Downsample2D) or isinstance(module, Upsample2D):
module.channels = int(module.channels * 0.75)
class Net(nn.Module):
def __init__(self, unet, decoder):
super().__init__()
del unet.time_embedding
new_conv_in = nn.Conv2d(16, 320, 3, padding=1)
new_conv_in.weight.data = unet.conv_in.weight.data.repeat(1, 4, 1, 1)
new_conv_in.bias.data = unet.conv_in.bias.data
unet.conv_in = new_conv_in
new_conv_out = nn.Conv2d(320, 342, 3, padding=1)
new_conv_out.weight.data = unet.conv_out.weight.data.repeat(86, 1, 1, 1)[:342]
new_conv_out.bias.data = unet.conv_out.bias.data.repeat(86,)[:342]
unet.conv_out = new_conv_out
def ResnetBlock2D_remove_time_emb_proj(module):
if isinstance(module, ResnetBlock2D):
del module.time_emb_proj
unet.apply(ResnetBlock2D_remove_time_emb_proj)
def BasicTransformerBlock_remove_cross_attn(module):
if isinstance(module, BasicTransformerBlock):
del module.attn2, module.norm2
unet.apply(BasicTransformerBlock_remove_cross_attn)
def set_inplace_to_true(module):
if isinstance(module, nn.Dropout) or isinstance(module, nn.SiLU):
module.inplace = True
unet.apply(set_inplace_to_true)
def replace_forward_methods(module):
if isinstance(module, CrossAttnDownBlock2D):
module.forward = types.MethodType(MyCrossAttnDownBlock2D_SD_forward, module)
elif isinstance(module, DownBlock2D):
module.forward = types.MethodType(MyDownBlock2D_SD_forward, module)
elif isinstance(module, UNetMidBlock2DCrossAttn):
module.forward = types.MethodType(MyUNetMidBlock2DCrossAttn_SD_forward, module)
elif isinstance(module, UpBlock2D):
module.forward = types.MethodType(MyUpBlock2D_SD_forward, module)
elif isinstance(module, CrossAttnUpBlock2D):
module.forward = types.MethodType(MyCrossAttnUpBlock2D_SD_forward, module)
elif isinstance(module, ResnetBlock2D):
module.forward = types.MethodType(MyResnetBlock2D_SD_forward, module)
elif isinstance(module, Transformer2DModel):
module.forward = types.MethodType(MyTransformer2DModel_SD_forward, module)
unet.apply(replace_forward_methods)
unet.forward = types.MethodType(MyUNet2DConditionModel_SD_forward, unet)
halve_channels(unet)
unet.body = nn.Sequential(
*unet.down_blocks,
unet.mid_block,
*unet.up_blocks,
unet.conv_norm_out,
unet.conv_act,
unet.conv_out,
)
del decoder.conv_in, decoder.up_blocks, decoder.conv_norm_out, decoder.conv_act, decoder.conv_out
self.body = nn.Sequential(
nn.PixelUnshuffle(2),
unet,
decoder.mid_block,
)
def forward(self, x):
return self.body(x) |