复数空间下2D数组向量归一化:补全fft_normalize函数实现
复数域向量归一化实现(支持2D数组)
我有一组按行存储在2D数组中的向量,生成要求为:
MEAN=0, STD-DEV=1/SQRT(vec_len)
操作前后需保持该归一化形式,需要在复数空间中实现归一化,已有函数框架如下:
@staticmethod def fft_normalize(x, dim=DEF_DIM): cx = rfft(x, dim=dim) .... rv = irfft(cx_proj, dim=dim) return rv
请补全中间代码
现有实值归一化函数
@staticmethod def normalize(a, dim=DEF_DIM): norm=torch.linalg.norm(a,dim=dim) # if torch.eq(norm,0) : return torch.divide(a,st.MIN) if dim is not None : norm = norm.unsqueeze(dim) return torch.divide(a,norm)
实值归一化测试示例
In [70]: st.normalize(x + 3) Out[70]: ([[0.05, 0.04, 0.05, ..., 0.04, 0.04, 0.04], [0.04, 0.04, 0.05, ..., 0.05, 0.04, 0.05], [0.05, 0.04, 0.05, ..., 0.04, 0.05, 0.04]]) In [71]: st.normalize(x + 5) Out[71]: ([[0.05, 0.04, 0.05, ..., 0.04, 0.04, 0.04], [0.04, 0.04, 0.05, ..., 0.05, 0.04, 0.04], [0.05, 0.04, 0.04, ..., 0.04, 0.05, 0.04]]) In [73]: st.normalize(x + 5).len() Out[73]: ([1.00, 1.00, 1.00]) In [74]: st.normalize(x + 3).len() Out[74]: ([1., 1., 1.]) In [75]: st.normalize(x).len() Out[75]: ([1.00, 1.00, 1.00]) #bad, need normalization In [76]: (x + 3).len() Out[76]: ([67.13, 67.13, 67.13]) @staticmethod def len(a,dim=DEF_DIM): return torch.linalg.norm(a,dim=dim)
个人尝试(仅支持1D向量)
@staticmethod def fft_normalize(x, dim=DEF_DIM):# Normalize a vector x in complex domain. c = rfft(x,dim=dim) ri = torch.vstack([c.real, c.imag]) norm = torch.abs(c) print(norm.shape, ri.shape) # norm = torch.linalg.norm(ri, dim=dim) # if dim is not None : norm = norm.unsqueeze(dim) if torch.any(torch.eq(norm,0)): norm[torch.eq(norm,0)] = st.MIN #!fixme ri= torch.divide(ri,norm) #2D fails here c_proj = ri[0,:] + 1j * ri[1,:] rv = irfft(c_proj, dim=dim) return rv
改编方案(仅支持1D向量)
@staticmethod def fft_normalize(x, dim=DEF_DIM, dot_dim=None):# Normalize a vector x in complex domain. c = rfftn(x,dim=dim) c_conj = torch.conj(c) if dot_dim is None : dot_dim = st.dot_dims(c, c_conj) c_norm = torch.sqrt(torch.tensordot(c, c_conj, dims=dot_dim)) c_proj = torch.divide(c, c_norm) rv = irfftn(c_proj, dim=dim) return rv
解决方案(支持2D数组)
补全后的代码可支持2D数组按行处理,核心保证频域范数计算的维度匹配,避免广播错误:
import torch from torch.fft import rfft, irfft @staticmethod def fft_normalize(x, dim=DEF_DIM): cx = rfft(x, dim=dim) # 计算频域复数向量的L2范数 c_norm = torch.linalg.norm(cx, dim=dim) # 处理零范数,避免除零错误 c_norm = torch.where(c_norm == 0, torch.tensor(st.MIN, device=c_norm.device), c_norm) # 扩展维度以匹配cx的形状,支持广播除法 if dim is not None: c_norm = c_norm.unsqueeze(dim) # 频域归一化 cx_proj = cx / c_norm # 逆FFT转回实域 rv = irfft(cx_proj, dim=dim) return rv
关键说明:
- 范数计算:直接使用
torch.linalg.norm可正确计算复数向量的L2范数,无需拆分实部虚部手动求和。 - 零范数处理:用
torch.where实现批量条件替换,比单值判断更高效且支持张量操作。 - 维度兼容:通过
unsqueeze(dim)扩展范数维度,确保和频域张量cx形状匹配,实现逐行的广播除法,完美支持2D数组。 - 归一化保持:频域归一化后逆变换回实域,输出向量仍满足MEAN=0、STD-DEV=1/SQRT(vec_len)的要求。
内容的提问来源于stack exchange,提问作者sten
相关产品推荐
相关产品推荐

