Coverage for wrapper/model_timing.py: 67%
86 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1import time
2import math
3import torch
5from typing import List
6from typing import Dict
7from typing import Optional
8from typing import Any
11class TimePeriod:
12 def __init__(self) -> None:
13 self.start_time: float = time.time()
14 self.end_time: Optional[float] = None
16 def end(self) -> None:
17 if torch.cuda.is_available():
18 torch.cuda.synchronize() # Ensure all prior CUDA ops are done
19 self.end_time = time.time()
21 def get_seconds(self) -> float:
22 if self.end_time is None:
23 return -1.0
24 return self.end_time - self.start_time
26 def __str__(self) -> str:
27 return f"{self.get_seconds():.2f}"
29 # TypeError: Object of type TimePeriod is not JSON serializable
30 def to_dict(self) -> Dict[str, float]:
31 return {
32 "start_time": self.start_time,
33 "end_time": self.end_time or -1.0,
34 "duration_seconds": round(self.get_seconds(), 2)
35 }
38class Timer:
39 def __init__(self) -> None:
40 self.timing = {}
41 self.timing["total"] = TimePeriod()
43 def start(self, event_name: str = "total") -> None:
44 self.timing[event_name] = TimePeriod()
46 def end(self, event_name: str = "total") -> None:
47 if event_name not in self.timing:
48 raise ValueError(f"Event {event_name} not found in {self.timing.keys()}.")
49 self.timing[event_name].end()
51 def get_last_event_name(self) -> Optional[str]:
52 if not self.timing:
53 return None
54 # Assumes that keys are added in order
55 return list(self.timing.keys())[-1]
57 def get_total_seconds(self) -> float:
58 if "total" not in self.timing:
59 return -1.0
60 return self.timing["total"].get_seconds()
62 def __str__(self) -> str:
63 str_ret = ""
64 for key, val in self.timing.items():
65 str_ret += f"{key}: {val.get_seconds():.3f}, "
66 return str_ret[:-2]
68 def to_dict(self) -> dict:
69 return {k: round(v.get_seconds(), 2) for k, v in self.timing.items()}
71 def to_timestamps(
72 self,
73 group: Optional[str] = None,
74 subgroup: Optional[str] = None
75 ) -> List[Dict[str, Any]]:
76 events = []
77 for key, val in self.timing.items():
78 id_key = f"{group}_{key}" if group else key
79 if subgroup:
80 id_key = f"{subgroup}_{key}"
81 event = {
82 "id": id_key,
83 "content": key,
84 # ceil and floor to nearest ms to avoid overlap, keep in seconds
85 "start": math.ceil(val.start_time * 1000) / 1000 if val.start_time is not None else None,
86 "end": math.floor(val.end_time * 1000) / 1000 if val.end_time is not None else None,
87 "duration_seconds": val.get_seconds()
88 }
89 if group:
90 event["group"] = group
91 if subgroup:
92 event["subgroup"] = subgroup
93 event["className"] = subgroup
94 events.append(event)
95 return events
98class LoadTimer(Timer):
99 def __init__(self) -> None:
100 super().__init__()
102 def __str__(self) -> str:
103 '''
104 For video generation:
105 text_encoder
106 image_encoder
107 vae
108 dit
109 '''
110 if "text_encoder" not in self.timing:
111 return ""
113 return f"{self.timing['text_encoder'].get_seconds():.3f}," + \
114 "{self.timing['image_encoder'].get_seconds():.3f}," + \
115 "{self.timing['vae'].get_seconds():.3f}," + \
116 "{self.timing['dit'].get_seconds():.3f}," + \
117 "{self.get_total_seconds():.3f}"
120class GenTimer(Timer):
121 def __init__(self) -> None:
122 super().__init__()
124 def __str__(self) -> str:
125 '''
126 For video/image generation, the order is:
127 text_encoder
128 image_encoder
129 vae_encoder
130 scheduler_setup
131 dit_{it}
132 dit_{it}_{it}
133 scheduler_{it}
134 vae_decoder
135 video_generation
136 '''
137 if "text_encoder" not in self.timing:
138 return ""
140 str_ret = f"{self.timing['text_encoder'].get_seconds():.3f}," + \
141 "{self.timing['image_encoder'].get_seconds():.3f}," + \
142 "{self.timing['vae_encoder'].get_seconds():.3f}," + \
143 "{self.timing['scheduler_setup'].get_seconds():.3f},"
144 dit_time = 0.0
145 scheduler_time = 0.0
146 for key, val in self.timing.items():
147 if key.startswith("dit_"):
148 dit_time += val.get_seconds()
149 elif key.startswith("scheduler_"):
150 scheduler_time += val.get_seconds()
151 str_ret += f"{dit_time:.3f},{scheduler_time:.3f},"
152 str_ret += f"{self.timing['vae_decoder'].get_seconds():.3f},{self.get_total_seconds():.3f}"
153 return str_ret