From e8265b603df37192f7e5d1cc61db7f5e03881665 Mon Sep 17 00:00:00 2001 From: Wei Sun Date: Tue, 7 Apr 2026 18:40:57 -0700 Subject: [PATCH] Fix TraceEvent.__init__() got an unexpected keyword argument pool_id MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Summary: ## Problem Newer PyTorch memory snapshots (introduced in D97221533) include a `pool_id` field in trace events. Mosaic's `TraceEvent.from_raw()` classmethod splats the raw dict directly into `TraceEvent(**raw_modified)`, causing a `TypeError` for any unrecognized field: ``` TypeError: TraceEvent.__init__() got an unexpected keyword argument 'pool_id' ``` This breaks all memory snapshot analysis tools (peak_memory_analysis, categorical_profiling, annotation_analysis, memory_diff) when used with snapshots from newer PyTorch versions. ## Root Cause `TraceEvent.from_raw()` in `fbcode/mosaic/libmosaic/utils/data_utils.py` individually pops known extra fields (e.g., `frames`, `user_metadata`) before splatting the remaining dict into the dataclass constructor. When PyTorch added `pool_id` to snapshot trace events (D97221533), this field leaked through as an unexpected kwarg. ## Fix Replaced the individual `.pop()` approach (previously used for `user_metadata` in D88310416) with **generic field filtering** — `raw_modified` is now filtered to only include keys that match actual `TraceEvent` dataclass fields using `dataclasses.fields()`. This is future-proof: any new fields PyTorch adds to snapshot trace events will be silently ignored without requiring another code change. ## Changes 1. **Updated import** (line 13): Added `fields as dataclass_fields` to the `dataclasses` import 2. **Generic field filtering** (lines 158-166): Replaced individual `.pop()` calls with `{k: v for k, v in raw_modified.items() if k in valid_fields}` filtering Reviewed By: basilwong Differential Revision: D99884123 --- mosaic/libmosaic/utils/data_utils.py | 15 ++++++++++++--- 1 file changed, 12 insertions(+), 3 deletions(-) diff --git a/mosaic/libmosaic/utils/data_utils.py b/mosaic/libmosaic/utils/data_utils.py index a73abd7..c925cec 100644 --- a/mosaic/libmosaic/utils/data_utils.py +++ b/mosaic/libmosaic/utils/data_utils.py @@ -10,7 +10,7 @@ import enum import logging from collections import defaultdict -from dataclasses import dataclass, field +from dataclasses import dataclass, field, fields as dataclass_fields from typing import Any, List, Optional, Union NUM_GPUS_PER_HOST = 8 @@ -154,10 +154,19 @@ def from_raw( ) del raw_modified["frames"] - raw_modified.pop("user_metadata", None) + + # Filter to only known TraceEvent fields to prevent future breakage + # from new fields added to PyTorch snapshots (e.g. pool_id, user_metadata) + known_fields = {f.name for f in dataclass_fields(cls)} + explicitly_set = {"classification", "custom_category", "annotation"} + filtered = { + k: v + for k, v in raw_modified.items() + if k in known_fields and k not in explicitly_set + } return cls( - **raw_modified, + **filtered, classification=classification, custom_category=custom_category or "unknown", annotation=annotation,