如何对图像执行2D小波散射?代码报错求助
问题分析与解决方案
错误原因
- 数据类型不匹配:你将图像转换为
float64类型的numpy数组,再转为torch的float64张量,但Kymatio的小波滤波器默认使用float32类型,导致逐元素乘法(cdgmm)操作时类型冲突。 - 前端与输入类型混用:你指定了
frontend='numpy',但传入的是torch张量,导致内部处理逻辑混乱。
修正方案
以下提供两种可行的修正代码,任选其一即可:
方案1:使用PyTorch前端(推荐)
import torch from kymatio import Scattering2D import numpy as np from PIL import Image FILENAME = "./square.png" image = Image.open(FILENAME).convert("L") # 转换为float32类型(匹配Kymatio默认滤波器类型) a = np.array(image).astype(np.float32) x = torch.from_numpy(a) # Scattering2D要求输入为4维张量:(batch, channels, height, width) x = x.unsqueeze(0).unsqueeze(0) imageSize = x.shape[2:] print(imageSize) # 使用默认torch前端(或显式指定frontend='torch') scattering = Scattering2D(J=2, shape=imageSize, L=8) Sx = scattering(x) print(Sx.size())
方案2:使用Numpy前端
from kymatio import Scattering2D import numpy as np from PIL import Image FILENAME = "./square.png" image = Image.open(FILENAME).convert("L") # 转换为float32类型 a = np.array(image).astype(np.float32) # 增加batch和channel维度,满足输入格式要求 a = a[np.newaxis, np.newaxis, :, :] imageSize = a.shape[2:] print(imageSize) # 指定numpy前端,传入numpy数组 scattering = Scattering2D(J=2, shape=imageSize, frontend='numpy', L=8) Sx = scattering(a) print(Sx.shape)
关键注意点
- 输入数据类型必须与Kymatio内部滤波器类型一致,优先使用
float32(PyTorch和Kymatio的默认类型)。 - Scattering2D的输入必须是4维结构:
(批量大小, 通道数, 高度, 宽度),需手动补充前两个维度。 - 前端类型要与输入数据类型匹配:torch张量对应torch前端,numpy数组对应numpy前端,不可混用。
内容的提问来源于stack exchange,提问作者Peter Minev
相关产品推荐
相关产品推荐

