You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.19 05:01:00