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

PyTorch Temporal Fusion Transformer分块训练内存泄漏排查求助

问题描述

在Databricks平台(128GB内存、32核机器)上,基于pytorch_forecasting的TemporalFusionTransformer模型处理海量零售数据。数据以Parquet格式按30个数据ID分组(chunk)存储在Blob存储中,采用分块训练方式,但训练时内存持续增长,最终因OOM错误(进程退出码137)崩溃,已尝试手动清理内存仍未解决,附上训练代码寻求排查帮助:

from datetime import date
from typing import Dict, Any, List

from pytorch_forecasting import TemporalFusionTransformer, TimeSeriesDataSet
from pytorch_forecasting.data import GroupNormalizer
from pytorch_forecasting.metrics import QuantileLoss
from lightning.pytorch.callbacks import (
    ModelCheckpoint,
    EarlyStopping,
    LearningRateMonitor

)
import torch
from lightning.pytorch.trainer import Trainer
import pandas as pd
import gc
import sys
import tracemalloc


from src.spark_session import get_spark_session
from src.modelisation.utils import (
    inverse_transform_scalers,
)
from conf.config import (
    target_column,
    columns_ids,
    MODELISATION_TFT,
    CATEGORICAL_ENCODERS,
    TFT_MODEL,
    num_workers_to_use,
)


def train_model_by_chunks(params: Dict[str, Any], list_chunks_path: List[str]):
    tft = None
    print("Starting iterate over chunks")

    # Load encoders
    categorical_label_encoders = torch.load("/dbfs" + CATEGORICAL_ENCODERS.as_posix())
    nchunks_alread_trained = 0

    # Create the trainer

    early_stop_callback = EarlyStopping(
        monitor="val_loss", min_delta=1e-4, patience=10, verbose=False, mode="min"
    )
    lr_logger = LearningRateMonitor()
    model_checkpoint = ModelCheckpoint(
        dirpath=("/dbfs" + TFT_MODEL.as_posix()),
        filename="best-checkpoint-tft",
        save_top_k=1,
        verbose=True,
        monitor="val_loss",
        mode="min",
        enable_version_counter=False,
    )

    trainer = Trainer(
        max_epochs=params["max_epochs"],
        accelerator="cpu",
        enable_model_summary=True,
        gradient_clip_val=0.1,
        callbacks=[model_checkpoint],
        default_root_dir=("/dbfs" + MODELISATION_TFT.as_posix()),
    )

    # Démarrer le suivi de la mémoire
    tracemalloc.start()

    for j, chunk_path in enumerate(list_chunks_path[nchunks_alread_trained:]):
        i = j + nchunks_alread_trained
        print(
            f"Dealing with chunk {i+1} : {i*params['chunk_size_train']} eans have been treated."
        )
        chunk = pd.read_parquet(chunk_path)
        chunk = preparation_timeseries_dataset_training_chunk(
            chunk_df=chunk, params=params
        )

        # Train model
        training = TimeSeriesDataSet(
            chunk,
            time_idx="time_index",
            target=target_column,
            group_ids=columns_ids,
            static_categoricals=params["static_categoricals"],
            static_reals=params["static_reals"],
            time_varying_known_categoricals=params["time_varying_known_categoricals"],
            time_varying_known_reals=params["time_varying_known_reals"],
            time_varying_unknown_reals=params["time_varying_unknown_reals"],
            target_normalizer=GroupNormalizer(
                groups=columns_ids, transformation="softplus"
            ),
            add_relative_time_idx=True,
            add_target_scales=True,
            max_prediction_length=params["max_prediction_length"],
            max_encoder_length=params["max_encoder_length"],
            categorical_encoders=categorical_label_encoders,
        )

        validation = TimeSeriesDataSet.from_dataset(
            training, chunk, predict=True, stop_randomization=True
        )

        train_dataloader = training.to_dataloader(
            train=True, batch_size=params["batch_size"], num_workers=num_workers_to_use
        )
        val_dataloader = validation.to_dataloader(
            train=False,
            batch_size=params["batch_size"] * 10,
            num_workers=num_workers_to_use,
        )

        if i == 0:
            # Initialize the model for the first chunk
            tft = TemporalFusionTransformer.from_dataset(
                training,
                learning_rate=0.03,
                hidden_size=16,
                attention_head_size=2,
                dropout=0.1,
                hidden_continuous_size=8,
                loss=QuantileLoss(),
                log_interval=-1,
                reduce_on_plateau_patience=4,
            )
        else:
            # Load the model for subsequent chunks
            tft = TemporalFusionTransformer.load_from_checkpoint(
                "/dbfs" + (TFT_MODEL / "best-checkpoint-tft.ckpt").as_posix()
            )
            print(
                f"Model loaded -- Number of parameters in network: {tft.size()/1e3:.1f}k"
            )

        trainer.fit(
            tft, train_dataloaders=train_dataloader, val_dataloaders=val_dataloader
        )

        # Capturer un instantané avant la suppression des objets
        snapshot_before = tracemalloc.take_snapshot()

        top_stats_before = snapshot_before.statistics("lineno")

        print("[ Top 10 Before deleting]")
        for stat in top_stats_before[:10]:
            print(stat)

        del chunk
        del train_dataloader
        del val_dataloader
        del training
        del validation
        del tft

        # Delete unused
        print(
            f"Before gc collect - Nombre total d'objets en mémoire: {len(gc.get_objects())}"
        )
        print(
            f"Before gc collect - Somme des objets en mémoire : {sum(sys.getsizeof(obj) for obj in gc.get_objects())*10e-9} GB"
        )
        gc.collect()

        print(
            f"After gc collect - Nombre total d'objets en mémoire: {len(gc.get_objects())}"
        )
        print(
            f"After gc collect - Somme des objets en mémoire : {sum(sys.getsizeof(obj) for obj in gc.get_objects())*10e-9} GB"
        )

        # Compare memoire après la suppresson des objets
        snapshot_after = tracemalloc.take_snapshot()
        stats = snapshot_after.compare_to(snapshot_before, "lineno")
        print("[ Top 10 différences ]")
        for stat in stats[:10]:
            print(stat)
    tracemalloc.stop()
排查与优化方案

1. 解决Trainer与Callback的状态累积问题

当前代码在循环外创建Trainer及回调函数,多次fit后会累积训练历史、日志、回调状态(如EarlyStopping的patience计数),不仅导致内存泄漏,还会干扰训练逻辑。需将Trainer和回调的创建移至循环内,每次处理新chunk时重置训练状态:

  • 每次循环重新初始化EarlyStopping、ModelCheckpoint等回调,避免状态残留
  • 每次循环创建新的Trainer实例,防止旧实例持有模型、数据的引用

2. 彻底释放DataLoader与Dataset内存

  • 关闭子进程泄漏:num_workers开启的子进程可能未被正确销毁,可先设置num_workers=0排查;若确认是子进程问题,保持persistent_workers=False(默认)确保每次epoch后销毁workers
  • 强制释放torch缓存:即使在CPU环境下,torch也会保留部分内存缓存,在gc.collect()后添加torch.empty_cache()强制释放
  • 减少不必要列加载:读取Parquet时指定columns参数,仅加载训练所需字段,降低单chunk内存占用

3. 优化模型加载与销毁流程

  • 解绑Trainer与模型:在删除模型前,调用trainer.clear_callbacks()、trainer.reset_train_loop()清空Trainer的关联状态,避免持有模型参数、优化器的引用
  • 选择性加载checkpoint:加载模型时,若不需要复用之前的优化器状态,可指定strict=False并忽略优化器加载,减少内存占用:
    tft = TemporalFusionTransformer.load_from_checkpoint(
        checkpoint_path,
        strict=False,
        optimizer=None,
        loss=QuantileLoss()
    )
    

4. 精准定位内存泄漏点

调整tracemalloc的使用逻辑,聚焦单个chunk处理周期内的内存变化:

  • 在循环开始前(读取chunk前)拍摄快照
  • 在循环结束后(gc.collect()后)拍摄快照
  • 对比两次快照,定位数据读取、Dataset创建、模型训练等环节的内存泄漏点

5. Databricks平台特定优化

  • 调整进程内存限制:通过Spark UI查看驱动进程内存使用情况,在集群配置中调整驱动内存阈值
  • 优化Parquet读取:使用storage_options={"cache_type": "none"}禁用Blob存储的IO缓存,避免额外内存占用
优化后的代码示例
from datetime import date
from typing import Dict, Any, List

from pytorch_forecasting import TemporalFusionTransformer, TimeSeriesDataSet
from pytorch_forecasting.data import GroupNormalizer
from pytorch_forecasting.metrics import QuantileLoss
from lightning.pytorch.callbacks import (
    ModelCheckpoint,
    EarlyStopping,
    LearningRateMonitor

)
import torch
from lightning.pytorch.trainer import Trainer
import pandas as pd
import gc
import sys
import tracemalloc


from src.spark_session import get_spark_session
from src.modelisation.utils import (
    inverse_transform_scalers,
)
from conf.config import (
    target_column,
    columns_ids,
    MODELISATION_TFT,
    CATEGORICAL_ENCODERS,
    TFT_MODEL,
    num_workers_to_use,
)


def train_model_by_chunks(params: Dict[str, Any], list_chunks_path: List[str]):
    print("Starting iterate over chunks")

    # Load encoders
    categorical_label_encoders = torch.load("/dbfs" + CATEGORICAL_ENCODERS.as_posix())
    nchunks_alread_trained = 0

    # Démarrer le suivi de la mémoire
    tracemalloc.start()

    for j, chunk_path in enumerate(list_chunks_path[nchunks_alread_trained:]):
        i = j + nchunks_alread_trained
        print(
            f"Dealing with chunk {i+1} : {i*params['chunk_size_train']} eans have been treated."
        )
        
        # 循环内创建Trainer与回调,避免状态累积
        early_stop_callback = EarlyStopping(
            monitor="val_loss", min_delta=1e-4, patience=10, verbose=False, mode="min"
        )
        lr_logger = LearningRateMonitor()
        model_checkpoint = ModelCheckpoint(
            dirpath=("/dbfs" + TFT_MODEL.as_posix()),
            filename="best-checkpoint-tft",
            save_top_k=1,
            verbose=True,
            monitor="val_loss",
            mode="min",
            enable_version_counter=False,
        )

        trainer = Trainer(
            max_epochs=params["max_epochs"],
            accelerator="cpu",
            enable_model_summary=True,
            gradient_clip_val=0.1,
            callbacks=[model_checkpoint, early_stop_callback, lr_logger],
            default_root_dir=("/dbfs" + MODELISATION_TFT.as_posix()),
        )
        
        # 读取前拍快照
        snapshot_start = tracemalloc.take_snapshot()
        
        # 仅读取所需列,减少内存
        required_columns = ["time_index", target_column] + columns_ids + params["static_categoricals"] + params["static_reals"] + params["time_varying_known_categoricals"] + params["time_varying_known_reals"] + params["time_varying_unknown_reals"]
        chunk = pd.read_parquet(chunk_path, columns=required_columns)
        chunk = preparation_timeseries_dataset_training_chunk(chunk_df=chunk, params=params)

        # Train model
        training = TimeSeriesDataSet(
            chunk,
            time_idx="time_index",
            target=target_column,
            group_ids=columns_ids,
            static_categoricals=params["static_categoricals"],
            static_reals=params["static_reals"],
            time_varying_known_categoricals=params["time_varying_known_categoricals"],
            time_varying_known_reals=params["time_varying_known_reals"],
            time_varying_unknown_reals=params["time_varying_unknown_reals"],
            target_normalizer=GroupNormalizer(
                groups=columns_ids, transformation="softplus"
            ),
            add_relative_time_idx=True,
            add_target_scales=True,
            max_prediction_length=params["max_prediction_length"],
            max_encoder_length=params["max_encoder_length"],
            categorical_encoders=categorical_label_encoders,
        )

        validation = TimeSeriesDataSet.from_dataset(
            training, chunk, predict=True, stop_randomization=True
        )

        train_dataloader = training.to_dataloader(
            train=True, batch_size=params["batch_size"], num_workers=num_workers_to_use, persistent_workers=False
        )
        val_dataloader = validation.to_dataloader(
            train=False,
            batch_size=params["batch_size"] * 10,
            num_workers=num_workers_to_use, persistent_workers=False
        )

        if i == 0:
            # Initialize the model for the first chunk
            tft = TemporalFusionTransformer.from_dataset(
                training,
                learning_rate=0.03,
                hidden_size=16,
                attention_head_size=2,
                dropout=0.1,
                hidden_continuous_size=8,
                loss=QuantileLoss(),
                log_interval=-1,
                reduce_on_plateau_patience=4,
            )
        else:
            # Load the model without optimizer state to save memory
            tft = TemporalFusionTransformer.load_from_checkpoint(
                "/dbfs" + (TFT_MODEL / "best-checkpoint-tft.ckpt").as_posix(),
                strict=False,
                optimizer=None,
                loss=QuantileLoss()
            )
            print(
                f"Model loaded -- Number of parameters in network: {tft.size()/1e3:.1f}k"
            )

        trainer.fit(
            tft, train_dataloaders=train_dataloader, val_dataloaders=val_dataloader
        )

        # 解绑Trainer与模型,清空状态
        trainer.clear_callbacks()
        trainer.reset_train_loop()

        # Capturer un instantané avant la suppression des objets
        snapshot_before = tracemalloc.take_snapshot()

        top_stats_before = snapshot_before.statistics("lineno")

        print("[ Top 10 Before deleting]")
        for stat in top_stats_before[:10]:
            print(stat)

        # 按依赖顺序删除,避免引用残留
        del tft
        del train_dataloader
        del val_dataloader
        del validation
        del training
        del chunk

        # 强制释放torch缓存与GC
        torch.empty_cache()
        print(
            f"Before gc collect - Nombre total d'objets en mémoire: {len(gc.get_objects())}"
        )
        print(
            f"Before gc collect - Somme des objets en mémoire : {sum(sys.getsizeof(obj) for obj in gc.get_objects())*10e-9} GB"
        )
        gc.collect()
        torch.empty_cache()

        print(
            f"After gc collect - Nombre total d'objets en mémoire: {len(gc.get_objects())}"
        )
        print(
            f"After gc collect - Somme des objets en mémoire : {sum(sys.getsizeof(obj) for obj in gc.get_objects())*10e-9} GB"
        )

        # 对比整个循环的内存变化
        snapshot_end = tracemalloc.take_snapshot()
        stats = snapshot_end.compare_to(snapshot_start, "lineno")
        print("[ Chunk Processing Memory Delta ]")
        for stat in stats[:10]:
            print(stat)
            
        # 删除当前循环的Trainer与回调
        del trainer
        del model_checkpoint
        del early_stop_callback
        del lr_logger
        gc.collect()
        torch.empty_cache()
        
    tracemalloc.stop()

内容的提问来源于stack exchange,提问作者Robin Vandamme

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 15:47:01