You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

复数空间下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

关键说明:

  1. 范数计算:直接使用torch.linalg.norm可正确计算复数向量的L2范数,无需拆分实部虚部手动求和。
  2. 零范数处理:用torch.where实现批量条件替换,比单值判断更高效且支持张量操作。
  3. 维度兼容:通过unsqueeze(dim)扩展范数维度,确保和频域张量cx形状匹配,实现逐行的广播除法,完美支持2D数组。
  4. 归一化保持:频域归一化后逆变换回实域,输出向量仍满足MEAN=0、STD-DEV=1/SQRT(vec_len)的要求。

内容的提问来源于stack exchange,提问作者sten

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.02 20:05:31