使用graph_weather库遇RuntimeError:张量维度不匹配问题求助
解决GraphWeatherForecaster张量维度不匹配问题
问题原因
报错提示输入特征维度(78)与模型输出维度(102)不匹配,核心原因是GraphWeatherForecaster类默认的输入/输出特征维度并非78(默认值对应官方数据集的特征数量),而你手动生成的输入特征为78维,导致模型输出维度和输入维度无法匹配。
解决方案
初始化GraphWeatherForecaster时,显式指定input_dim和output_dim参数为你的输入特征维度78,确保模型的输入输出维度与数据维度完全对齐。
修正后的完整代码:
import torch from graph_weather import GraphWeatherForecaster from graph_weather.models.losses import NormalizedMSELoss lat_lons = [] for lat in range(-90, 90, 1): for lon in range(0, 360, 1): lat_lons.append((lat, lon)) # 显式指定input_dim和output_dim为78,匹配自定义特征维度 model = GraphWeatherForecaster(lat_lons, input_dim=78, output_dim=78) features = torch.randn((2, len(lat_lons), 78)) out = model(features) criterion = NormalizedMSELoss(lat_lons=lat_lons, feature_variance=torch.randn((78,))) loss = criterion(out, features) loss.backward()
补充说明
- 若后续更换特征维度,只需同步修改
input_dim、output_dim以及输入特征张量的最后一维大小即可; - 你代码中
NormalizedMSELoss的feature_variance参数维度已设置为78维,这部分与特征维度匹配,无需调整。
内容的提问来源于stack exchange,提问作者Gianni Spear
相关产品推荐
相关产品推荐

