如何为TemporalFusionTransformer模型自定义命名?
解决方案:为TemporalFusionTransformer自定义模型名称
由于PyTorch Forecasting的TemporalFusionTransformer类本身没有内置的模型命名属性,你可以通过以下两种简单方式实现自定义命名:
方法1:直接为模型实例添加自定义属性
训练完成后,直接给模型对象赋值一个自定义名称属性,后续生成文件名时直接调用该属性即可:
from pytorch_forecasting import TemporalFusionTransformer from datetime import datetime # 训练模型流程 tft = TemporalFusionTransformer.from_dataset(your_dataset) tft.fit(your_train_dataloader) # 自定义模型名称 tft.name = "tft_sales_forecast_v1" # 生成指定格式的文件名 current_time = datetime.now().strftime("%Y%m%d_%H%M%S") log_filename = f"{tft.name}@{current_time}.txt" # 写入评估指标到文件 with open(log_filename, "w") as f: f.write(f"MAE: {your_mae_score}\n") f.write(f"RMSE: {your_rmse_score}\n")
方法2:继承扩展TemporalFusionTransformer类
如果希望在初始化模型时就指定名称,可以通过继承原类扩展一个带命名参数的子类,更符合工程规范:
from pytorch_forecasting import TemporalFusionTransformer from datetime import datetime class NamedTFT(TemporalFusionTransformer): def __init__(self, model_name: str, *args, **kwargs): super().__init__(*args, **kwargs) self.model_name = model_name # 初始化带名称的模型 tft = NamedTFT( model_name="tft_user_behavior_v2", dataset=your_dataset, hidden_size=64, attention_head_size=4, # 其他模型参数... ) tft.fit(your_train_dataloader) # 生成文件名并写入指标 current_time = datetime.now().strftime("%Y%m%d_%H%M%S") log_filename = f"{tft.model_name}@{current_time}.txt" with open(log_filename, "w") as f: f.write(f"评估指标内容...")
另外,恳请声望大于1500的用户在标签栏添加temporalfusiontransformer标签,我因声望不足无法添加新标签。
内容的提问来源于stack exchange,提问作者ptushev
相关产品推荐
相关产品推荐

