Ouzhang's picture
Add files using upload-large-folder tool
31dc8dc verified
Raw
History Blame Contribute Delete
5.47 kB
from __future__ import annotations
import csv
import json
import os
import time
from contextlib import nullcontext
from pathlib import Path
from typing import Any
import torch
def profile_enabled() -> bool:
return bool(os.getenv("DIFFULEX_PROFILE_DIR"))
def record_function(name: str):
if not profile_enabled():
return nullcontext()
return torch.profiler.record_function(name)
def _env_bool(name: str, default: bool) -> bool:
raw = os.getenv(name)
if raw is None:
return default
return raw.strip().lower() in {"1", "true", "yes", "on"}
def _env_int(name: str, default: int) -> int:
raw = os.getenv(name)
if raw is None or raw.strip() == "":
return default
return int(raw)
class TorchProfileSession:
def __init__(self, component: str, *, rank: int | None = None):
self.enabled = profile_enabled()
self.component = component
self.rank = rank
self.prof: torch.profiler.profile | None = None
self.started = False
self.stopped = False
self.steps = 0
self.active_steps = _env_int("DIFFULEX_PROFILE_ACTIVE_STEPS", 200)
def start(self) -> None:
if not self.enabled or self.started or self.stopped:
return
activities = [torch.profiler.ProfilerActivity.CPU]
if torch.cuda.is_available():
activities.append(torch.profiler.ProfilerActivity.CUDA)
kwargs = dict(
activities=activities,
record_shapes=_env_bool("DIFFULEX_PROFILE_RECORD_SHAPES", False),
profile_memory=_env_bool("DIFFULEX_PROFILE_MEMORY", False),
with_stack=_env_bool("DIFFULEX_PROFILE_WITH_STACK", False),
with_modules=_env_bool("DIFFULEX_PROFILE_WITH_MODULES", False),
acc_events=True,
)
try:
self.prof = torch.profiler.profile(**kwargs)
except TypeError:
kwargs.pop("acc_events", None)
self.prof = torch.profiler.profile(**kwargs)
self.prof.start()
self.started = True
def step(self) -> None:
if not self.enabled or self.stopped:
return
self.start()
self.steps += 1
if self.prof is not None and not self.stopped:
self.prof.step()
if self.active_steps > 0 and self.steps >= self.active_steps:
self.stop()
def stop(self) -> None:
if not self.enabled or not self.started or self.stopped:
return
assert self.prof is not None
self.stopped = True
try:
if torch.cuda.is_available():
torch.cuda.synchronize()
self.prof.stop()
self._export()
except Exception as exc:
prefix = self._prefix()
prefix.with_suffix(".error.txt").write_text(f"{type(exc).__name__}: {exc}\n", encoding="utf-8")
def _prefix(self) -> Path:
root = Path(os.environ["DIFFULEX_PROFILE_DIR"]).expanduser().resolve()
root.mkdir(parents=True, exist_ok=True)
rank_part = f".rank{self.rank}" if self.rank is not None else ""
stamp = os.getenv("DIFFULEX_PROFILE_RUN_ID") or time.strftime("%Y%m%d_%H%M%S")
return root / f"{stamp}.{self.component}{rank_part}"
@staticmethod
def _event_to_row(event: Any) -> dict[str, Any]:
self_device_time = getattr(
event,
"self_cuda_time_total",
getattr(event, "self_device_time_total", 0.0),
)
device_time = getattr(
event,
"cuda_time_total",
getattr(event, "device_time_total", 0.0),
)
self_device_memory = getattr(
event,
"self_cuda_memory_usage",
getattr(event, "self_device_memory_usage", 0),
)
return {
"name": event.key,
"count": event.count,
"self_cpu_time_total_us": getattr(event, "self_cpu_time_total", 0.0),
"cpu_time_total_us": getattr(event, "cpu_time_total", 0.0),
"self_cuda_time_total_us": self_device_time,
"cuda_time_total_us": device_time,
"self_cpu_memory_usage": getattr(event, "self_cpu_memory_usage", 0),
"self_cuda_memory_usage": self_device_memory,
"input_shapes": str(getattr(event, "input_shapes", "")),
}
def _export(self) -> None:
assert self.prof is not None
prefix = self._prefix()
trace_path = prefix.with_suffix(".trace.json")
summary_txt = prefix.with_suffix(".summary.txt")
summary_csv = prefix.with_suffix(".summary.csv")
summary_json = prefix.with_suffix(".summary.json")
sort_by = "self_cuda_time_total" if torch.cuda.is_available() else "self_cpu_time_total"
events = self.prof.key_averages()
rows = [self._event_to_row(event) for event in events]
rows.sort(key=lambda row: row["self_cuda_time_total_us"] or row["self_cpu_time_total_us"], reverse=True)
self.prof.export_chrome_trace(str(trace_path))
summary_txt.write_text(events.table(sort_by=sort_by, row_limit=200), encoding="utf-8")
with summary_csv.open("w", encoding="utf-8", newline="") as f:
writer = csv.DictWriter(f, fieldnames=list(rows[0]) if rows else ["name"])
writer.writeheader()
writer.writerows(rows)
summary_json.write_text(json.dumps(rows, indent=2), encoding="utf-8")