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

1import time 

2import math 

3import torch 

4 

5from typing import List 

6from typing import Dict 

7from typing import Optional 

8from typing import Any 

9 

10 

11class TimePeriod: 

12 def __init__(self) -> None: 

13 self.start_time: float = time.time() 

14 self.end_time: Optional[float] = None 

15 

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() 

20 

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 

25 

26 def __str__(self) -> str: 

27 return f"{self.get_seconds():.2f}" 

28 

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 } 

36 

37 

38class Timer: 

39 def __init__(self) -> None: 

40 self.timing = {} 

41 self.timing["total"] = TimePeriod() 

42 

43 def start(self, event_name: str = "total") -> None: 

44 self.timing[event_name] = TimePeriod() 

45 

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() 

50 

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] 

56 

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() 

61 

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] 

67 

68 def to_dict(self) -> dict: 

69 return {k: round(v.get_seconds(), 2) for k, v in self.timing.items()} 

70 

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 

96 

97 

98class LoadTimer(Timer): 

99 def __init__(self) -> None: 

100 super().__init__() 

101 

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 "" 

112 

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}" 

118 

119 

120class GenTimer(Timer): 

121 def __init__(self) -> None: 

122 super().__init__() 

123 

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 "" 

139 

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