File size: 7,208 Bytes
9f818c5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: OpenMDW-1.1

import contextlib
import os
import time

import torch

from cosmos_framework.utils import distributed, log
from cosmos_framework.utils.easy_io import easy_io

# (qsh 2024-11-23)  credits
# https://github.com/pytorch/torchtitan/blob/main/torchtitan/profiling.py

# how much memory allocation/free ops to record in memory snapshots
MEMORY_SNAPSHOT_MAX_ENTRIES = 100000


@contextlib.contextmanager
def maybe_enable_profiling(config, *, global_step: int = 0):
    # get user defined profiler settings
    enable_profiling = config.trainer.profiling.enable_profiling
    profile_freq = config.trainer.profiling.profile_freq

    if enable_profiling:
        trace_dir = os.path.join(config.job.path_local, "torch_trace")
        if distributed.get_rank() == 0:
            os.makedirs(trace_dir, exist_ok=True)

        rank = distributed.get_rank()

        def trace_handler(prof):
            curr_trace_dir_name = "iteration_" + str(prof.step_num)
            curr_trace_dir = os.path.join(trace_dir, curr_trace_dir_name)
            if not os.path.exists(curr_trace_dir):
                os.makedirs(curr_trace_dir, exist_ok=True)

            log.info(f"Dumping traces at step {prof.step_num}")
            begin = time.monotonic()
            if rank in config.trainer.profiling.target_ranks:
                prof.export_chrome_trace(f"{curr_trace_dir}/rank{rank}_trace.json.gz")
            log.info(f"Finished dumping traces in {time.monotonic() - begin:.2f} seconds")

        log.info(f"Profiling active. Traces will be saved at {trace_dir}")

        if not os.path.exists(trace_dir):
            os.makedirs(trace_dir, exist_ok=True)

        warmup, active = config.trainer.profiling.profile_warmup, 1
        wait = profile_freq - (active + warmup)
        assert wait >= 0, "profile_freq must be greater than or equal to warmup + active"

        with torch.profiler.profile(
            activities=[
                torch.profiler.ProfilerActivity.CPU,
                torch.profiler.ProfilerActivity.CUDA,
            ],
            schedule=torch.profiler.schedule(wait=wait, warmup=warmup, active=active),
            on_trace_ready=trace_handler,
            record_shapes=config.trainer.profiling.record_shape,
            profile_memory=config.trainer.profiling.profile_memory,
            with_stack=config.trainer.profiling.with_stack,
            with_modules=config.trainer.profiling.with_modules,
        ) as torch_profiler:
            torch_profiler.step_num = global_step
            yield torch_profiler
    else:
        torch_profiler = contextlib.nullcontext()
        yield None


@contextlib.contextmanager
def maybe_enable_memory_snapshot(config, *, global_step: int = 0):
    enable_snapshot = config.trainer.profiling.enable_memory_snapshot
    if enable_snapshot:
        if config.trainer.profiling.save_s3:
            snapshot_dir = "s3://rundir"
        else:
            snapshot_dir = os.path.join(config.job.path_local, "memory_snapshot")
            if distributed.get_rank() == 0:
                os.makedirs(snapshot_dir, exist_ok=True)

        rank = torch.distributed.get_rank()

        class MemoryProfiler:
            def __init__(self, step_num: int, freq: int):
                torch.cuda.memory._record_memory_history(max_entries=MEMORY_SNAPSHOT_MAX_ENTRIES)
                # when resume training, we start from the last step
                self.step_num = step_num
                self.freq = freq

            def step(self, exit_ctx: bool = False):
                self.step_num += 1
                if not exit_ctx and self.step_num % self.freq != 0:
                    return
                if not exit_ctx:
                    curr_step = self.step_num
                    dir_name = f"iteration_{curr_step}"
                else:
                    # dump as iteration_0_exit if OOM at iter 1
                    curr_step = self.step_num - 1
                    dir_name = f"iteration_{curr_step}_exit"
                curr_snapshot_dir = os.path.join(snapshot_dir, dir_name)
                if not config.trainer.profiling.save_s3 and not os.path.exists(curr_snapshot_dir):
                    os.makedirs(curr_snapshot_dir, exist_ok=True)
                log.info(f"Dumping memory snapshot at step {curr_step}")
                begin = time.monotonic()

                if rank in config.trainer.profiling.target_ranks:
                    easy_io.dump(
                        torch.cuda.memory._snapshot(),
                        f"{curr_snapshot_dir}/rank{rank}_memory_snapshot.pickle",
                    )
                log.info(f"Finished dumping memory snapshot in {time.monotonic() - begin:.2f} seconds")

        log.info(f"Memory profiler active. Snapshot will be saved at {snapshot_dir}")
        profiler = MemoryProfiler(global_step, config.trainer.profiling.profile_freq)
        try:
            yield profiler
        except torch.cuda.OutOfMemoryError as e:
            profiler.step(exit_ctx=True)
    else:
        yield None


@contextlib.contextmanager
def maybe_enable_nsys_profiling(config, *, global_step: int = 0):
    """Context manager for Nsight Systems profiling via cudaProfilerStart/Stop.

    Usage: launch training with
        nsys profile --capture-range=cudaProfilerApi --capture-range-end=stop python ...
    and set trainer.profiling.enable_nsys=true, profile_freq=<iter>.

    Reuses the torch-profile flags (profile_freq, target_ranks, profile_warmup).
    The profiler is started `profile_warmup` iterations before the target and
    stopped right after it.
    """
    enable_nsys = config.trainer.profiling.enable_nsys
    if not enable_nsys:
        yield None
        return

    rank = distributed.get_rank()
    target_ranks = config.trainer.profiling.target_ranks
    freq = config.trainer.profiling.profile_freq
    warmup = config.trainer.profiling.profile_warmup

    active_iter = freq - 1  # profile_freq=5001 profiles iter 5000
    start_iter = max(0, active_iter - warmup)

    class NsysProfiler:
        def __init__(self, step_num: int):
            self.step_num = step_num
            self._profiling = False

        def step(self):
            self.step_num += 1
            if rank not in target_ranks:
                return
            if self.step_num == start_iter and not self._profiling:
                log.info(f"[Nsys] Starting CUDA profiler at iter {self.step_num} (active iter: {active_iter})")
                torch.cuda.cudart().cudaProfilerStart()
                self._profiling = True
            if self.step_num == active_iter + 1 and self._profiling:
                torch.cuda.cudart().cudaProfilerStop()
                self._profiling = False
                log.info(f"[Nsys] Stopped CUDA profiler at iter {self.step_num}")

    log.info(f"[Nsys] Profiling enabled. Will capture iter {start_iter}-{active_iter} on ranks {target_ranks}")
    profiler = NsysProfiler(global_step)
    try:
        yield profiler
    finally:
        if profiler._profiling:
            torch.cuda.cudart().cudaProfilerStop()