Files
change_hair_3090/hair_service_sd/faceseg/tma.py
T
colomi 0eb61f3e60 初始化换发型项目:3个微服务代码 + 部署脚本
包含:
- 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 排除,
由网盘单独上传。
2026-07-11 18:11:49 +08:00

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