如何覆盖FastAI为权重文件路径添加models/前缀与.pth后缀的行为?
解决FastAI自动拼接权重文件路径的问题
在fastai==2.7.7版本中,调用model.load()加载权重时,框架会自动给路径添加models/前缀和.pth后缀,导致传入完整路径后出现重复拼接的错误(比如models/data/models/my_weights.pth.pth)。以下是几种覆盖该行为的解决办法:
方法1:使用绝对路径或Path对象
FastAI的load方法仅对相对字符串路径做自动拼接,传入绝对路径字符串或者pathlib.Path对象时,会直接使用原路径,不会添加额外前缀后缀。代码示例:
from pathlib import Path def inference(model_params: ModelParams, device: str, dataloader: DataLoader, id: str): model = load_learner(model_params.model, cpu=='cpu') # 方式1:使用Path对象传入完整路径 weight_path = Path('data/models/my_weights.pth') model.load(weight_path) # 方式2:直接传入绝对路径字符串 # model.load('/home/me/project/data/models/my_weights.pth')
方法2:修改Learner的默认路径
Learner默认会在自身path属性指定的目录下的models子目录找权重文件,修改path为权重所在目录后,只需传入文件名(无需后缀)即可:
from pathlib import Path def inference(model_params: ModelParams, device: str, dataloader: DataLoader, id: str): model = load_learner(model_params.model, cpu=='cpu') # 将Learner的默认路径设置为权重文件所在目录 model.path = Path('data/models') # 仅传入文件名,框架不会额外拼接路径 model.load('my_weights')
方法3:手动用torch.load加载权重
绕过FastAI的load方法,直接使用PyTorch原生API加载权重并挂载到模型上,完全控制路径逻辑:
import torch def inference(model_params: ModelParams, device: str, dataloader: DataLoader, id: str): model = load_learner(model_params.model, cpu=='cpu') # 直接加载权重文件 state_dict = torch.load('data/models/my_weights.pth', map_location=device) # 将权重加载到模型中 model.model.load_state_dict(state_dict)
内容的提问来源于stack exchange,提问作者DanielBell99
相关产品推荐
相关产品推荐

