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
相关产品推荐
相关产品推荐

