Timestamp direct memory profile output
This commit is contained in:
parent
bc0cff88ce
commit
c3d72d9f1e
3 changed files with 10 additions and 3 deletions
|
|
@ -32,7 +32,7 @@ class H3Checkpoint:
|
||||||
with fastsafe_open(filenames=[self.path], nogds=True, device=self._fast_device()) as checkpoint:
|
with fastsafe_open(filenames=[self.path], nogds=True, device=self._fast_device()) as checkpoint:
|
||||||
self._no_mmap_tensors = {
|
self._no_mmap_tensors = {
|
||||||
name: checkpoint.get_tensor(name).clone().detach()
|
name: checkpoint.get_tensor(name).clone().detach()
|
||||||
for name in checkpoint.get_keys()
|
for name in checkpoint.keys()
|
||||||
}
|
}
|
||||||
return self._no_mmap_tensors
|
return self._no_mmap_tensors
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -172,7 +172,7 @@ class MiniMaxH3VideoVAE(nn.Module):
|
||||||
with fastsafe_open(filenames=[str(path)], nogds=True, device=fast_device) as checkpoint:
|
with fastsafe_open(filenames=[str(path)], nogds=True, device=fast_device) as checkpoint:
|
||||||
available_weights = {
|
available_weights = {
|
||||||
name: checkpoint.get_tensor(name).clone().detach()
|
name: checkpoint.get_tensor(name).clone().detach()
|
||||||
for name in checkpoint.get_keys()
|
for name in checkpoint.keys()
|
||||||
}
|
}
|
||||||
available = set(available_weights)
|
available = set(available_weights)
|
||||||
missing = sorted(set(expected) - available)
|
missing = sorted(set(expected) - available)
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,8 @@ import argparse
|
||||||
import os
|
import os
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import subprocess
|
import subprocess
|
||||||
|
import time
|
||||||
|
from datetime import datetime, timezone
|
||||||
import warnings
|
import warnings
|
||||||
|
|
||||||
warnings.filterwarnings("ignore", message="Found GPU0 NVIDIA GB10 which is of cuda capability 12.1.*", category=UserWarning)
|
warnings.filterwarnings("ignore", message="Found GPU0 NVIDIA GB10 which is of cuda capability 12.1.*", category=UserWarning)
|
||||||
|
|
@ -32,11 +34,15 @@ parser.add_argument("--attention", choices=("sage2", "sdpa", "sage3"), default="
|
||||||
parser.add_argument("--model-timesteps-capture", type=Path, help="Directory containing captured input_XX.pt H3 timesteps for strict parity checks.")
|
parser.add_argument("--model-timesteps-capture", type=Path, help="Directory containing captured input_XX.pt H3 timesteps for strict parity checks.")
|
||||||
parser.add_argument("--profile-memory", action="store_true")
|
parser.add_argument("--profile-memory", action="store_true")
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
started = time.perf_counter()
|
||||||
|
last_report = started
|
||||||
|
|
||||||
|
|
||||||
def report_memory(stage: str) -> None:
|
def report_memory(stage: str) -> None:
|
||||||
|
global last_report
|
||||||
if not args.profile_memory:
|
if not args.profile_memory:
|
||||||
return
|
return
|
||||||
|
now = time.perf_counter()
|
||||||
rss_kb = 0
|
rss_kb = 0
|
||||||
try:
|
try:
|
||||||
with open("/proc/self/status", encoding="utf-8") as file:
|
with open("/proc/self/status", encoding="utf-8") as file:
|
||||||
|
|
@ -48,7 +54,8 @@ def report_memory(stage: str) -> None:
|
||||||
pass
|
pass
|
||||||
cuda_alloc = torch.cuda.memory_allocated() / 1024**3 if torch.cuda.is_available() else 0.0
|
cuda_alloc = torch.cuda.memory_allocated() / 1024**3 if torch.cuda.is_available() else 0.0
|
||||||
cuda_reserved = torch.cuda.memory_reserved() / 1024**3 if torch.cuda.is_available() else 0.0
|
cuda_reserved = torch.cuda.memory_reserved() / 1024**3 if torch.cuda.is_available() else 0.0
|
||||||
print({"stage": stage, "rss_gb": round(rss_kb / 1024**2, 3), "cuda_alloc_gb": round(cuda_alloc, 3), "cuda_reserved_gb": round(cuda_reserved, 3), "fast_safetensors": os.getenv("H3_FAST_SAFETENSORS", ""), "disable_mmap": os.getenv("H3_DISABLE_MMAP", "")}, flush=True)
|
print({"stage": stage, "ts": datetime.now(timezone.utc).isoformat(), "epoch_s": round(time.time(), 3), "elapsed_s": round(now - started, 3), "delta_s": round(now - last_report, 3), "rss_gb": round(rss_kb / 1024**2, 3), "cuda_alloc_gb": round(cuda_alloc, 3), "cuda_reserved_gb": round(cuda_reserved, 3), "fast_safetensors": os.getenv("H3_FAST_SAFETENSORS", ""), "disable_mmap": os.getenv("H3_DISABLE_MMAP", "")}, flush=True)
|
||||||
|
last_report = now
|
||||||
|
|
||||||
report_memory("start")
|
report_memory("start")
|
||||||
checkpoint = H3Checkpoint("/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
|
checkpoint = H3Checkpoint("/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue