forked from ccf-ai-infra/GPUCodeForces
155 lines
5.0 KiB
Python
155 lines
5.0 KiB
Python
import torch
|
||
import torch.nn as nn
|
||
import torch.nn.functional as F
|
||
|
||
# 定义维度常量
|
||
N, C, H, W = 32, 64, 56, 56
|
||
EPS = 1e-6
|
||
|
||
|
||
class EvoNormS0(nn.Module):
|
||
"""
|
||
EvoNorm-S0: Evolving Normalization-Activation Layers (Sample-based, no batch dependency)
|
||
|
||
公式:
|
||
v = Var(x) = mean(x^2) - mean(x)^2
|
||
y = x / sqrt(v + eps) * gamma + beta
|
||
y = y * sigmoid(x * w)
|
||
|
||
其中 gamma, beta, w 是可学习参数
|
||
"""
|
||
|
||
def __init__(self, num_channels, eps, nonlinear=True):
|
||
super().__init__()
|
||
self.eps = eps
|
||
self.nonlinear = nonlinear # 是否使用非线性激活
|
||
|
||
# 可学习的缩放和偏移参数(类似 BatchNorm)
|
||
self.gamma = nn.Parameter(torch.ones(1, num_channels, 1, 1))
|
||
self.beta = nn.Parameter(torch.zeros(1, num_channels, 1, 1))
|
||
|
||
# 非线性门控参数
|
||
if self.nonlinear:
|
||
self.v = nn.Parameter(torch.ones(1, num_channels, 1, 1))
|
||
|
||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||
# 1. 计算实例级方差
|
||
# var = E[x^2] - E[x]^2
|
||
x_sq_mean = torch.mean(x * x, dim=[2, 3], keepdim=True)
|
||
x_mean = torch.mean(x, dim=[2, 3], keepdim=True)
|
||
var = x_sq_mean - x_mean * x_mean
|
||
|
||
# 2. 归一化
|
||
x_normalized = x / torch.sqrt(var + self.eps)
|
||
|
||
# 3. 仿射变换
|
||
y = x_normalized * self.gamma + self.beta
|
||
|
||
# 4. 非线性门控(可选)
|
||
if self.nonlinear:
|
||
y = y * torch.sigmoid(x * self.v)
|
||
|
||
return y
|
||
|
||
|
||
class EvoNormB0(nn.Module):
|
||
"""
|
||
EvoNorm-B0: Evolving Normalization-Activation Layers (Batch-based)
|
||
|
||
公式:
|
||
Instance Norm: x_in = (x - mean(x)) / sqrt(var(x) + eps)
|
||
Batch Norm stats: rolling_var = momentum * rolling_var + (1-momentum) * batch_var
|
||
y = x_in * gamma + beta
|
||
y = y * sigmoid(x * w)
|
||
"""
|
||
|
||
def __init__(self, num_channels, eps, momentum=0.1, nonlinear=True):
|
||
super().__init__()
|
||
self.eps = eps
|
||
self.momentum = momentum
|
||
self.nonlinear = nonlinear
|
||
|
||
# 可学习参数
|
||
self.gamma = nn.Parameter(torch.ones(1, num_channels, 1, 1))
|
||
self.beta = nn.Parameter(torch.zeros(1, num_channels, 1, 1))
|
||
|
||
# 非线性门控参数
|
||
if self.nonlinear:
|
||
self.v = nn.Parameter(torch.ones(1, num_channels, 1, 1))
|
||
|
||
# 运行时统计量(用于推理)
|
||
self.register_buffer('running_var', torch.ones(1, num_channels, 1, 1))
|
||
self.register_buffer('num_batches_tracked', torch.tensor(0, dtype=torch.long))
|
||
|
||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||
if self.training:
|
||
# 训练模式:计算当前批次的统计量
|
||
# 1. 实例归一化
|
||
x_mean = torch.mean(x, dim=[2, 3], keepdim=True)
|
||
x_var = torch.var(x, dim=[2, 3], keepdim=True, unbiased=False)
|
||
|
||
# 2. 更新运行统计量(跨批次的方差)
|
||
batch_var = torch.mean(x_var, dim=0, keepdim=True)
|
||
with torch.no_grad():
|
||
self.running_var = (1 - self.momentum) * self.running_var + self.momentum * batch_var
|
||
self.num_batches_tracked += 1
|
||
|
||
# 3. 归一化
|
||
x_normalized = (x - x_mean) / torch.sqrt(x_var + self.eps)
|
||
else:
|
||
# 推理模式:使用运行统计量
|
||
x_mean = torch.mean(x, dim=[2, 3], keepdim=True)
|
||
x_normalized = (x - x_mean) / torch.sqrt(self.running_var + self.eps)
|
||
|
||
# 4. 仿射变换
|
||
y = x_normalized * self.gamma + self.beta
|
||
|
||
# 5. 非线性门控
|
||
if self.nonlinear:
|
||
y = y * torch.sigmoid(x * self.v)
|
||
|
||
return y
|
||
|
||
|
||
class Model(nn.Module):
|
||
"""
|
||
EvoNorm 模型包装器
|
||
默认使用 EvoNorm-S0(无批次依赖,更适合小批量)
|
||
"""
|
||
|
||
def __init__(self, evonorm_gamma, evonorm_beta, evonorm_v=None, use_b0=False):
|
||
super().__init__()
|
||
|
||
# 选择 EvoNorm 变体
|
||
if use_b0:
|
||
self.evonorm = EvoNormB0(C, EPS, nonlinear=(evonorm_v is not None))
|
||
else:
|
||
self.evonorm = EvoNormS0(C, EPS, nonlinear=(evonorm_v is not None))
|
||
|
||
# 初始化参数
|
||
with torch.no_grad():
|
||
self.evonorm.gamma.data.copy_(evonorm_gamma.view(1, C, 1, 1))
|
||
self.evonorm.beta.data.copy_(evonorm_beta.view(1, C, 1, 1))
|
||
|
||
if evonorm_v is not None and self.evonorm.nonlinear:
|
||
self.evonorm.v.data.copy_(evonorm_v.view(1, C, 1, 1))
|
||
|
||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||
return self.evonorm(x)
|
||
|
||
|
||
def get_inputs():
|
||
"""生成测试输入"""
|
||
x = torch.randn(N, C, H, W, dtype=torch.float32)
|
||
return [x]
|
||
|
||
|
||
def get_init_inputs():
|
||
"""
|
||
生成初始化参数
|
||
返回 [gamma, beta, v]
|
||
"""
|
||
evonorm_gamma = torch.ones(1, C, 1, 1)
|
||
evonorm_beta = torch.zeros(1, C, 1, 1)
|
||
evonorm_v = torch.ones(1, C, 1, 1) # 门控参数
|
||
return [evonorm_gamma, evonorm_beta, evonorm_v] |