Python序列化报错求助:无法pickle局部类对象'load_pt..spike'
解决pickle序列化自定义局部类的问题
问题根源
你在load_pt函数内部定义spike类,pickle无法序列化这种局部类——因为pickle需要类的定义存在于模块的全局作用域,才能在反序列化时找到它。哪怕加了global spike也没用,因为这个类是函数运行时动态生成的,模块加载阶段全局作用域里根本没有这个类的定义,反序列化时会找不到。
解决方案:把类移到全局作用域
最直接且适合新手的方法,是把spike类定义在模块的顶部(函数外面),然后在load_pt里创建类的实例并返回(你原来的代码返回的是类本身,注释里说要返回对象,应该是返回实例)。
修改后的代码
- 先在模块全局区域定义
spike类:
class spike: def __init__(self, values, select, chlabels, fs, soz): self.values = values self.select = select self.chlabels = chlabels self.fs = fs self.soz = soz
- 改写
load_pt函数,创建实例并返回:
def load_pt(ptname, data_directory): """ input: ptname, a string containing the name of the patient you want to load. example: 'ABC123' output: object: spike ---- contains: List of 1000 random spikes: spike.select and their subsequent: values: spike.values, chlabels: spike.chlabels, fs: spike.fs, soz channels: spike.soz """ val = mat73.loadmat(data_directory + '/values/values_{}.mat'.format(ptname)) val2 = val['values_all'] select_spikes = loadmat(data_directory + '/randi/randi_{}.mat'.format(ptname)) select_spikes = select_spikes['select_spikes'] ch_labels = loadmat(data_directory + '/chlabels/chlabels_{}.mat'.format(ptname)) ch_labels = ch_labels['ch_labels_all'] fs_all = loadmat(data_directory + '/fs/fs_{}.mat'.format(ptname)) fs_all = fs_all['fs_all'] SOZ_chlabels = pd.read_csv(data_directory + '/pt_data/SOZ_channels.csv') pt_all = pd.read_csv(data_directory + '/pt_data/ptname_all.csv') pt_name = ("'{}'".format(ptname)) whichpt = pt_all.index[pt_all['ptname'] == pt_name].tolist() clean_SOZ_chlabels = prep_clean_soz(SOZ_chlabels) # 创建spike实例并返回 return spike( values=val2, select=select_spikes, chlabels=ch_labels, fs=fs_all, soz=clean_SOZ_chlabels[whichpt[0]] )
为什么这样能解决问题
现在spike类存在于模块的全局作用域,pickle序列化实例时,只需要记录类的路径(模块名+类名),反序列化时就能直接找到类的定义,不会再报错。
内容的提问来源于stack exchange,提问作者Carlos Aguila
相关产品推荐
相关产品推荐

