Skip to content

Commit b3d7f85

Browse files
mengluy0125facebook-github-bot
authored andcommitted
Use the consolidated snapshot API in Unitrace to support Zoomer
Summary: Similar to D48210543. Update the training_hooks to use the Unitrace memory snapshot APIs. This allows us to maintain a singel path for memory snapshot APIs, and also collect important details such as snapshot location for Zoomer. Pull Request resolved: facebookresearch#613 Reviewed By: frabu6, jackiexu1992 Differential Revision: D48368150 Pulled By: HugeEngine fbshipit-source-id: ed5d819153bdebef4143fa0794855bfc0640ccda
1 parent 8d072eb commit b3d7f85

4 files changed

Lines changed: 29 additions & 126 deletions

File tree

d2go/runner/config_defaults.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -138,10 +138,10 @@ def _add_detectron2go_runner_default_cfg(_C: CN) -> None:
138138
# Profiler
139139
_C.PROFILERS = ["default_flop_counter"]
140140

141-
# GPU memory profiler
141+
# Snapshot memory profiling
142142
add_memory_profiler_configs(_C)
143143

144-
# Zoomer memory profiling
144+
# Zoomer Kineto memory profiling
145145
add_zoomer_default_config(_C)
146146

147147
# Checkpointing-specific config

d2go/runner/default_runner.py

Lines changed: 3 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111

1212
import detectron2.utils.comm as comm
1313
import torch
14+
from aiplatform.monitoring.unitrace.memory_snapshot import attach_oom_logger
1415
from d2go.checkpoint.api import is_distributed_checkpoint
1516
from d2go.checkpoint.fsdp_checkpoint import FSDPCheckpointer
1617
from d2go.config import CfgNode, CONFIG_SCALING_METHOD_REGISTRY, temp_defrost
@@ -38,11 +39,7 @@
3839
get_generalized_rcnn_runner_default_cfg,
3940
)
4041

41-
from d2go.runner.training_hooks import (
42-
D2GoGpuMemorySnapshot,
43-
TRAINER_HOOKS_REGISTRY,
44-
update_hooks_from_registry,
45-
)
42+
from d2go.runner.training_hooks import update_hooks_from_registry
4643
from d2go.trainer.fsdp import get_grad_scaler
4744
from d2go.trainer.helper import parse_precision_from_string
4845
from d2go.utils.abnormal_checker import (
@@ -51,7 +48,6 @@
5148
get_writers,
5249
)
5350
from d2go.utils.flop_calculator import attach_profilers
54-
from d2go.utils.gpu_memory_profiler import attach_oom_logger
5551
from d2go.utils.helper import D2Trainer, TensorboardXWriter
5652
from d2go.utils.misc import get_tensorboard_log_dir
5753
from d2go.utils.visualization import DataLoaderVisWrapper, VisualizationEvaluator
@@ -150,20 +146,6 @@ def default_scale_quantization_configs(cfg, new_world_size):
150146
)
151147

152148

153-
@TRAINER_HOOKS_REGISTRY.register()
154-
def add_memory_profiler_hook(hooks, cfg: CfgNode):
155-
# Add GPU memory snapshot profiler to diagnose GPU OOM issues and benchmark memory usage during model training
156-
if cfg.get("MEMORY_PROFILER", CfgNode()).get("ENABLED", False):
157-
hooks.append(
158-
D2GoGpuMemorySnapshot(
159-
cfg.OUTPUT_DIR,
160-
log_n_steps=cfg.MEMORY_PROFILER.LOG_N_STEPS,
161-
log_during_train_at=cfg.MEMORY_PROFILER.LOG_DURING_TRAIN_AT,
162-
trace_max_entries=cfg.MEMORY_PROFILER.TRACE_MAX_ENTRIES,
163-
)
164-
)
165-
166-
167149
@fb_overwritable()
168150
def prepare_fb_model(cfg: CfgNode, model: torch.nn.Module) -> torch.nn.Module:
169151
return model
@@ -345,9 +327,7 @@ def _build_model(self, cfg, eval_only=False):
345327
def build_model(self, cfg, eval_only=False):
346328
# Attach memory profiler to GPU OOM events
347329
if cfg.get("MEMORY_PROFILER", CfgNode()).get("ENABLED", False):
348-
attach_oom_logger(
349-
cfg.OUTPUT_DIR, trace_max_entries=cfg.MEMORY_PROFILER.TRACE_MAX_ENTRIES
350-
)
330+
attach_oom_logger(bucket="d2go_traces")
351331

352332
model = self._build_model(cfg, eval_only)
353333
model = prepare_fb_model(cfg, model)

d2go/runner/training_hooks.py

Lines changed: 24 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -3,12 +3,17 @@
33
import logging
44
from typing import List
55

6-
from d2go.config import CfgNode
6+
from aiplatform.monitoring.unitrace.memory_snapshot import (
7+
export_memory_snapshot,
8+
start_record_memory_history,
9+
stop_record_memory_history,
10+
)
711

8-
from d2go.utils.gpu_memory_profiler import log_memory_snapshot, record_memory_history
12+
from d2go.config import CfgNode
913

1014
from detectron2.engine.train_loop import HookBase
1115
from detectron2.utils.registry import Registry
16+
from mobile_cv.torch.utils_pytorch import comm
1217

1318

1419
logger = logging.getLogger(__name__)
@@ -41,29 +46,34 @@ class D2GoGpuMemorySnapshot(HookBase):
4146

4247
def __init__(
4348
self,
44-
output_dir,
4549
log_n_steps: int = 3,
4650
log_during_train_at: int = 550,
47-
trace_max_entries: int = 1000000,
51+
manifold_bucket: str = "d2go_traces",
52+
root_manifold_path: str = "tree/memory_snapshot",
4853
) -> None:
49-
self.output_dir = output_dir
50-
self.step = 0
5154
self.log_n_steps = log_n_steps
5255
self.log_during_train_at = log_during_train_at
53-
self.trace_max_entries = trace_max_entries
56+
self.manifold_bucket = manifold_bucket
57+
self.root_manifold_path = root_manifold_path
5458
logger.warning(
5559
"WARNING: Memory snapshot profiler is enabled. This may cause ranks to die and training jobs to get stuck. Please use with caution."
5660
)
5761

5862
def before_step(self):
5963
if self.trainer.iter == self.log_during_train_at:
60-
record_memory_history(self.trace_max_entries)
64+
logger.info(
65+
f"[itrn-{self.trainer.iter}] Starting memory snapshot recording"
66+
)
67+
start_record_memory_history()
6168

6269
def after_step(self):
63-
if self.step == self.log_n_steps - 1:
64-
log_memory_snapshot(self.output_dir, file_prefix=f"iter{self.trainer.iter}")
65-
6670
if self.trainer.iter == self.log_during_train_at + self.log_n_steps - 1:
67-
log_memory_snapshot(self.output_dir, file_prefix=f"iter{self.trainer.iter}")
68-
69-
self.step += 1
71+
export_memory_snapshot(
72+
worker_name=f"rank-{comm.get_rank()}",
73+
bucket=self.manifold_bucket,
74+
root_manifold_path=self.root_manifold_path,
75+
)
76+
logger.info(
77+
f"[itrn-{self.trainer.iter}] Stopping memory snapshot recording"
78+
)
79+
stop_record_memory_history()

d2go/utils/gpu_memory_profiler.py

Lines changed: 0 additions & 87 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,6 @@
11
import logging
2-
import os
3-
import pickle
42

5-
import torch
63
from d2go.config import CfgNode as CN
7-
from detectron2.utils.file_io import PathManager
8-
from mobile_cv.torch.utils_pytorch import comm
9-
from torch.cuda._memory_viz import segment_plot, trace_plot
104

115
logger: logging.Logger = logging.getLogger(__name__)
126

@@ -29,84 +23,3 @@ def add_zoomer_default_config(_C: CN):
2923
False # Do not enable by default, since it may cause performance regression
3024
)
3125
_C.ZOOMER.ENABLE_MEMORY_PROFILING = False
32-
33-
34-
def omm_logger_wrapper(output_dir):
35-
def oom_logger(
36-
device: int, alloc: int, device_alloc: int, device_free: int
37-
) -> None:
38-
"""
39-
Log memory snapshot in the event of CUDA OOM.
40-
"""
41-
logger.info(
42-
f"Saving memory snapshot device: {device}, alloc: {alloc}, device_alloc: {device_alloc}, device_free: {device_free}"
43-
)
44-
try:
45-
log_memory_snapshot(output_dir, file_prefix="oom")
46-
except Exception as e:
47-
logger.error(f"Failed to log memory snapshot during OOM {e}")
48-
49-
return oom_logger
50-
51-
52-
def log_memory_snapshot(output_dir: str, file_prefix: str = "") -> None:
53-
"""
54-
Log memory snapshots to output_dir
55-
"""
56-
if not torch.cuda.is_available():
57-
logger.info("CUDA unavailable. Not logging snapshot")
58-
return
59-
60-
try:
61-
rank = comm.get_rank()
62-
save_dir = os.path.join(
63-
output_dir, "memory_snapshot", f"{file_prefix}_rank{rank}"
64-
)
65-
logger.info(f"Logging memory snapshot to {save_dir}")
66-
snapshot = torch.cuda.memory._snapshot()
67-
dump_snapshot(save_dir, snapshot)
68-
except Exception as e:
69-
logger.error(f"Failed to log memory snapshot to {save_dir}: {e}")
70-
71-
72-
def dump_snapshot(save_dir: str, snapshot):
73-
"""
74-
Dump memory snapshot and useful plots to save_dir.
75-
This is a rewrite of torch.cuda.memory._dump_snapshot() with PathManager.
76-
"""
77-
if not PathManager.exists(save_dir):
78-
PathManager.mkdirs(save_dir)
79-
with PathManager.open(os.path.join(save_dir, "snapshot.pickle"), "wb") as f:
80-
pickle.dump(snapshot, f)
81-
with PathManager.open(os.path.join(save_dir, "trace_plot.html"), "w") as f:
82-
f.write(trace_plot(snapshot))
83-
with PathManager.open(os.path.join(save_dir, "segment_plot.html"), "w") as f:
84-
f.write(segment_plot(snapshot))
85-
logger.info(f"Saved memory snapshot to {save_dir}")
86-
87-
88-
def record_memory_history(trace_max_entries=1000000) -> None:
89-
"""
90-
Start recording memory history and stack traces.
91-
"""
92-
if not torch.cuda.is_available():
93-
logger.info("CUDA unavailable. Not recording memory history")
94-
return
95-
96-
torch.cuda.memory._record_memory_history(
97-
enabled="all", max_entries=trace_max_entries
98-
)
99-
logger.info("Started recording memory history")
100-
101-
102-
def attach_oom_logger(output_dir, trace_max_entries=1000000) -> None:
103-
"""
104-
Start recording memory history and attach the OOM logger.
105-
"""
106-
if not torch.cuda.is_available():
107-
logger.info("CUDA unavailable. Not attaching OOM logger")
108-
return
109-
110-
record_memory_history(trace_max_entries)
111-
torch._C._cuda_attach_out_of_memory_observer(omm_logger_wrapper(output_dir))
112-
logger.info("Attached GPU OOM logger")

0 commit comments

Comments
 (0)