PyTorch-Forecasting中quantile() dtype不匹配RuntimeError求助
环境信息
- PyTorch-Forecasting版本: 0.10.2
- PyTorch版本: 1.12.1
- Python版本: 3.10.4
- 操作系统: Windows
预期行为
无报错
实际行为
运行时触发如下错误:
File c:\Users\josepeeterson.er\Miniconda3\envs\pytorch\lib\site-packages\pytorch_forecasting\metrics\base_metrics.py:979, in DistributionLoss.to_quantiles(self, y_pred, quantiles, n_samples)
977 except NotImplementedError: # resort to derive quantiles empirically
978 samples = torch.sort(self.sample(y_pred, n_samples), -1).values
--> 979 quantiles = torch.quantile(samples, torch.tensor(quantiles, device=samples.device), dim=2).permute(1, 2, 0)
980 return quantilesRuntimeError: quantile() q tensor must be same dtype as the input tensor
该错误来自框架内部代码,无法直接修改相关逻辑,如何让两个张量的数据类型保持一致?未使用GPU。
输入数据为每4小时从参数(9,0.5)的负二项分布采样得到,其余时间值为0,目标是验证DeepAR模型能否学习该时序模式。
复现代码
from pytorch_forecasting.data.examples import generate_ar_data import matplotlib.pyplot as plt import pandas as pd from pytorch_forecasting.data import TimeSeriesDataSet from pytorch_forecasting.data import NaNLabelEncoder from pytorch_lightning.callbacks import EarlyStopping, LearningRateMonitor import pytorch_lightning as pl from pytorch_forecasting import NegativeBinomialDistributionLoss, DeepAR import torch from pytorch_forecasting.data.encoders import TorchNormalizer # 修正原代码中data被赋值为列表的错误 data = pd.read_csv('1_f_nbinom_train.csv') data["date"] = pd.Timestamp("2021-08-24") + pd.to_timedelta(data.time_idx, "H") data['_hour_of_day'] = data["date"].dt.hour.astype(str) data['_day_of_week'] = data["date"].dt.dayofweek.astype(str) data['_day_of_month'] = data["date"].dt.day.astype(str) data['_day_of_year'] = data["date"].dt.dayofyear.astype(str) # 修正weekofyear弃用问题 data['_week_of_year'] = data["date"].dt.isocalendar().week.astype(str) data['_month_of_year'] = data["date"].dt.month.astype(str) data['_year'] = data["date"].dt.year.astype(str) max_encoder_length = 60 max_prediction_length = 20 training_cutoff = data["time_idx"].max() - max_prediction_length training = TimeSeriesDataSet( data.iloc[0:-620], time_idx="time_idx", target="value", categorical_encoders={ "series": NaNLabelEncoder(add_nan=True).fit(data.series), "_hour_of_day": NaNLabelEncoder(add_nan=True).fit(data._hour_of_day), "_day_of_week": NaNLabelEncoder(add_nan=True).fit(data._day_of_week), "_day_of_month": NaNLabelEncoder(add_nan=True).fit(data._day_of_month), "_day_of_year": NaNLabelEncoder(add_nan=True).fit(data._day_of_year), "_week_of_year": NaNLabelEncoder(add_nan=True).fit(data._week_of_year), "_year": NaNLabelEncoder(add_nan=True).fit(data._year) }, group_ids=["series"], min_encoder_length=max_encoder_length, max_encoder_length=max_encoder_length, min_prediction_length=max_prediction_length, max_prediction_length=max_prediction_length, time_varying_unknown_reals=["value"], time_varying_known_categoricals=["_hour_of_day","_day_of_week","_day_of_month","_day_of_year","_week_of_year","_year" ], time_varying_known_reals=["time_idx"], add_relative_time_idx=False, randomize_length=None, scalers=[], target_normalizer=TorchNormalizer(method="identity", center=False, transformation=None) ) validation = TimeSeriesDataSet.from_dataset( training, data.iloc[-620:-420], stop_randomization=True, ) batch_size = 64 train_dataloader = training.to_dataloader(train=True, batch_size=batch_size, num_workers=8) val_dataloader = validation.to_dataloader(train=False, batch_size=batch_size, num_workers=8) # save datasets training.save("training.pkl") validation.save("validation.pkl") early_stop_callback = EarlyStopping(monitor="val_loss", min_delta=1e-4, patience=5, verbose=False, mode="min") lr_logger = LearningRateMonitor() trainer = pl.Trainer( max_epochs=10, gpus=0, gradient_clip_val=0.1, limit_train_batches=30, limit_val_batches=3, callbacks=[lr_logger, early_stop_callback], ) deepar = DeepAR.from_dataset( training, learning_rate=0.1, hidden_size=32, dropout=0.1, loss=NegativeBinomialDistributionLoss(), log_interval=10, log_val_interval=3, ) print(f"Number of parameters in network: {deepar.size()/1e3:.1f}k") torch.set_num_threads(10) trainer.fit( deepar, train_dataloaders=train_dataloader, val_dataloaders=val_dataloader, )
解决方法
1. 升级PyTorch-Forecasting版本
该问题是PyTorch-Forecasting 0.10.2的已知兼容性Bug,后续版本已修复,直接升级即可:
pip install --upgrade pytorch-forecasting
2. 临时修改框架代码
如果无法升级,找到报错文件base_metrics.py(路径:c:\Users\josepeeterson.er\Miniconda3\envs\pytorch\lib\site-packages\pytorch_forecasting\metrics\base_metrics.py),将第979行修改为:
quantiles = torch.quantile(samples, torch.tensor(quantiles, device=samples.device, dtype=samples.dtype), dim=2).permute(1, 2, 0)
强制让quantiles张量与samples使用相同的数据类型。
3. 全局指定默认 dtype
在代码开头添加以下代码,统一浮点类型:
torch.set_default_dtype(torch.float32)
内容的提问来源于stack exchange,提问作者Jose_Peeterson

