包含: - hair_service_sd: 换发型/换发色算法服务 (端口 8801) - photo_service: LoRA 训练调度服务 (端口 32678) - stable-diffusion-webui: SD WebUI 推理服务 (端口 57860) - kohya_ss_home: 训练环境代码 - meidaojia: 监控测试脚本 - setup.sh: 一键部署脚本 (conda环境恢复 + 配置生成 + 完整性检查) - start_all_services.sh: 启动3个服务 - configure.ini.template: 路径模板化 (BASE_DIR自动推导) - conda_envs/py310.yml: py310 环境定义 大文件 (weights/, models/, data/, conda_envs/*.tar.gz 等) 通过 .gitignore 排除, 由网盘单独上传。
186 lines
7.2 KiB
Python
186 lines
7.2 KiB
Python
import torch
|
|
import torch.nn.functional as F
|
|
from torch import nn
|
|
import numpy as np
|
|
|
|
class SequenceConv(nn.ModuleList):
|
|
"""Sequence conv module.
|
|
|
|
Args:
|
|
in_channels (int): input tensor channel.
|
|
out_channels (int): output tensor channel.
|
|
kernel_size (int): convolution kernel size.
|
|
sequence_num (int): sequence length.
|
|
conv_cfg (dict): convolution config dictionary.
|
|
norm_cfg (dict): normalization config dictionary.
|
|
act_cfg (dict): activation config dictionary.
|
|
"""
|
|
|
|
def __init__(self, in_channels, out_channels, kernel_size, sequence_num):
|
|
super(SequenceConv, self).__init__()
|
|
self.in_channels = in_channels
|
|
self.out_channels = out_channels
|
|
self.kernel_size = kernel_size
|
|
self.sequence_num = sequence_num
|
|
for _ in range(sequence_num):
|
|
self.append(
|
|
nn.Sequential(
|
|
nn.Conv2d(self.in_channels, self.out_channels, self.kernel_size, 1, self.kernel_size // 2, bias=False),
|
|
nn.BatchNorm2d(self.out_channels),
|
|
nn.ReLU()
|
|
)
|
|
)
|
|
|
|
def forward(self, sequence_imgs):
|
|
"""
|
|
|
|
Args:
|
|
sequence_imgs (Tensor): TxBxCxHxW
|
|
|
|
Returns:
|
|
sequence conv output: TxBxCxHxW
|
|
"""
|
|
sequence_outs = []
|
|
assert sequence_imgs.shape[0] == self.sequence_num
|
|
for i, sequence_conv in enumerate(self):
|
|
sequence_out = sequence_conv(sequence_imgs[i, ...])
|
|
sequence_out = sequence_out.unsqueeze(0)
|
|
sequence_outs.append(sequence_out)
|
|
|
|
sequence_outs = torch.cat(sequence_outs, dim=0) # TxBxCxHxW
|
|
return sequence_outs
|
|
|
|
class MemoryModule(nn.Module):
|
|
"""Memory read module.
|
|
Args:
|
|
|
|
"""
|
|
|
|
def __init__(self,
|
|
matmul_norm=False):
|
|
super(MemoryModule, self).__init__()
|
|
self.matmul_norm = matmul_norm
|
|
|
|
def forward(self, memory_keys, memory_values, query_key, query_value):
|
|
"""
|
|
Memory Module forward.
|
|
Args:
|
|
memory_keys (Tensor): memory keys tensor, shape: TxBxCxHxW
|
|
memory_values (Tensor): memory values tensor, shape: TxBxCxHxW
|
|
query_key (Tensor): query keys tensor, shape: BxCxHxW
|
|
query_value (Tensor): query values tensor, shape: BxCxHxW
|
|
|
|
Returns:
|
|
Concat query and memory tensor.
|
|
"""
|
|
sequence_num, batch_size, key_channels, height, width = memory_keys.shape
|
|
_, _, value_channels, _, _ = memory_values.shape
|
|
assert query_key.shape[1] == key_channels and query_value.shape[1] == value_channels
|
|
memory_keys = memory_keys.permute(1, 2, 0, 3, 4).contiguous() # BxCxTxHxW
|
|
memory_keys = memory_keys.view(batch_size, key_channels, sequence_num * height * width) # BxCxT*H*W
|
|
|
|
query_key = query_key.view(batch_size, key_channels, height * width).permute(0, 2, 1).contiguous() # BxH*WxCk
|
|
key_attention = torch.bmm(query_key, memory_keys) # BxH*WxT*H*W
|
|
if self.matmul_norm:
|
|
key_attention = (key_channels ** -.5) * key_attention
|
|
key_attention = F.softmax(key_attention, dim=-1) # BxH*WxT*H*W
|
|
|
|
memory_values = memory_values.permute(1, 2, 0, 3, 4).contiguous() # BxCxTxHxW
|
|
memory_values = memory_values.view(batch_size, value_channels, sequence_num * height * width)
|
|
memory_values = memory_values.permute(0, 2, 1).contiguous() # BxT*H*WxC
|
|
memory = torch.bmm(key_attention, memory_values) # BxH*WxC
|
|
memory = memory.permute(0, 2, 1).contiguous() # BxCxH*W
|
|
memory = memory.view(batch_size, value_channels, height, width) # BxCxHxW
|
|
|
|
query_memory = torch.cat([query_value, memory], dim=1)
|
|
return query_memory
|
|
#
|
|
# class TMAHead(nn.Module):
|
|
# """TMAHead decoder for video semantic segmentation."""
|
|
#
|
|
# def __init__(self, sequence_num, key_channels, value_channels, num_classes=2, dropout_ratio=0):
|
|
# super(TMAHead, self).__init__()
|
|
#
|
|
# self.sequence_num = sequence_num
|
|
# self.memory_key_conv = nn.Sequential(
|
|
# SequenceConv(self.in_channels, key_channels, 1, sequence_num),
|
|
# SequenceConv(key_channels, key_channels, 3, sequence_num)
|
|
# )
|
|
# self.memory_value_conv = nn.Sequential(
|
|
# SequenceConv(self.in_channels, value_channels, 1, sequence_num),
|
|
# SequenceConv(value_channels, value_channels, 3, sequence_num)
|
|
# )
|
|
# self.query_key_conv = nn.Sequential(
|
|
# nn.Sequential(
|
|
# nn.Conv2d(self.in_channels, key_channels, 1, 1, 0, bias=False),
|
|
# nn.BatchNorm2d(key_channels),
|
|
# nn.ReLU()
|
|
# ),
|
|
# nn.Sequential(
|
|
# nn.Conv2d(key_channels, key_channels, 3, 1, 1, bias=False),
|
|
# nn.BatchNorm2d(key_channels),
|
|
# nn.ReLU()
|
|
# ),
|
|
# )
|
|
#
|
|
# self.query_value_conv = nn.Sequential(
|
|
# nn.Sequential(
|
|
# nn.Conv2d(self.in_channels, value_channels, 1, 1, 0, bias=False),
|
|
# nn.BatchNorm2d(value_channels),
|
|
# nn.ReLU()
|
|
# ),
|
|
# nn.Sequential(
|
|
# nn.Conv2d(value_channels, value_channels, 3, 1, 1, bias=False),
|
|
# nn.BatchNorm2d(value_channels),
|
|
# nn.ReLU()
|
|
# ),
|
|
# )
|
|
# self.memory_module = MemoryModule(matmul_norm=False)
|
|
# self.bottleneck = nn.Sequential(
|
|
# nn.Conv2d(value_channels * 2, self.channels, 3, 1, 1, bias=False),
|
|
# nn.BatchNorm2d(value_channels),
|
|
# nn.ReLU()
|
|
# )
|
|
#
|
|
# self.conv_seg = nn.Conv2d(self.channels, num_classes, kernel_size=1)
|
|
# if dropout_ratio > 0:
|
|
# self.dropout = nn.Dropout2d(dropout_ratio)
|
|
# else:
|
|
# self.dropout = None
|
|
#
|
|
# def cls_seg(self, feat):
|
|
# """Classify each pixel."""
|
|
# if self.dropout is not None:
|
|
# feat = self.dropout(feat)
|
|
# output = self.conv_seg(feat)
|
|
# return output
|
|
#
|
|
# def forward(self, inputs, sequence_imgs):
|
|
# """
|
|
# Forward fuction.
|
|
# Args:
|
|
# inputs (list[Tensor]): backbone multi-level outputs.
|
|
# sequence_imgs (list[Tensor]): len(sequence_imgs) is equal to batch_size,
|
|
# each element is a Tensor with shape of TxCxHxW.
|
|
#
|
|
# Returns:
|
|
# decoder logits.
|
|
# """
|
|
# x = inputs
|
|
# sequence_imgs = [y.unsqueeze(0) for y in sequence_imgs] # T, BxCxHxW
|
|
# sequence_imgs = torch.cat(sequence_imgs, dim=0) # TxBxCxHxW
|
|
# sequence_num, batch_size, channels, height, width = sequence_imgs.shape
|
|
#
|
|
# assert sequence_num == self.sequence_num
|
|
# memory_keys = self.memory_key_conv(sequence_imgs)
|
|
# memory_values = self.memory_value_conv(sequence_imgs)
|
|
# query_key = self.query_key_conv(x) # BxCxHxW
|
|
# query_value = self.query_value_conv(x) # BxCxHxW
|
|
#
|
|
# # memory read
|
|
# output = self.memory_module(memory_keys, memory_values, query_key, query_value)
|
|
# output = self.bottleneck(output)
|
|
# output = self.cls_seg(output)
|
|
#
|
|
# return output
|