GPUCodeForces/S1/uucoco_#12/CoVariance_torch.py

65 lines
1.5 KiB
Python

import torch
import torch.nn as nn
N, C, H, W = 32, 64, 56, 56
EPS = 1e-8
class Covariance(nn.Module):
def __init__(self, eps=1e-8, unbiased=True):
super().__init__()
self.unbiased = unbiased
def forward(self, x, y):
# 扁平化 x 和 y 从 (N, C, H, W) 到 (N, D)
x_flat = x.reshape(x.size(0), -1)
y_flat = y.reshape(y.size(0), -1)
# 元素总数 D = C * H * W
D = x_flat.size(1)
# 1. 计算均值 (Mean)
x_mean = x_flat.mean(dim=1, keepdim=True)
y_mean = y_flat.mean(dim=1, keepdim=True)
# 2. 居中化 (Centering)
x_centered = x_flat - x_mean
y_centered = y_flat - y_mean
# 3. 计算协方差和 (Covariance Sum)
# 结果维度: (N,)
cov_sum = (x_centered * y_centered).sum(dim=1)
# 4. 计算协方差 (除以 D 或 D-1)
if self.unbiased:
# 样本协方差
divisor = D - 1
else:
# 总体协方差
divisor = D
if divisor <= 0:
return cov_sum * 0.0
return cov_sum / divisor
class Model(nn.Module):
def __init__(self, unbiased=True):
super().__init__()
self.op = Covariance(unbiased=unbiased)
def forward(self, x, y):
return self.op(x, y)
def get_inputs():
x = torch.randn(N, C, H, W, dtype=torch.float32)
y = torch.randn(N, C, H, W, dtype=torch.float32)
return [x, y]
def get_init_inputs():
return []