如何在PySpark中使用Pandas UDF结合Scipy计算Welch PSD
错误原因
- 输入信号提取错误:GROUPED_MAP模式下传入UDF的
df包含signal、id两列,直接调用df.values会把无意义的id列数值带入计算,完全污染了输入信号,应提取df['signal'].values作为计算输入。 - 频率轴构造错误:
signal.welch返回的第一个结果f就是校准完成的全范围频率轴,无需手动重新生成。你实现的手动构造逻辑仅生成了从0到第一个频率间隔的数值,频率范围完全错误。 - 函数调用参数不一致:参考代码显式指定了
nperseg=1024,PySpark代码中未传该参数,scipy默认nperseg=256,参数不同计算结果自然存在差异。 - 绘图方式不一致:参考代码使用对数纵坐标
plt.semilogy展示结果,你使用线性坐标plt.plot,即便计算正确视觉效果也会有明显区别。
修正后核心代码
@pandas_udf(StructType([StructField('frequency',FloatType()), StructField('power_dens',FloatType())]), PandasUDFType.GROUPED_MAP) def welch(df): # 仅提取信号列,补充和参考代码一致的nperseg参数 f, Pxx_den = signal.welch(df['signal'].values, fs, nperseg=1024) # 直接使用返回的频率轴,无需手动构造 y = pd.DataFrame({'frequency':f.astype(float), 'power_dens':Pxx_den.astype(float)}) return y # 绘图时和参考逻辑保持一致 try_df = trying.toPandas() plt.figure(figsize=(15,10)) plt.semilogy(try_df['frequency'],try_df['power_dens']) plt.ylim([0.5e-3, 1]) plt.xlabel('frequency [Hz]') plt.ylabel('PSD [V**2/Hz]') plt.show()
内容的提问来源于stack exchange,提问作者Gilang Samudra Kuswana
相关产品推荐
相关产品推荐

