Skip to content

Commit 6cee324

Browse files
fix bali pid and bali result dict keys finding
1 parent 939b393 commit 6cee324

3 files changed

Lines changed: 22 additions & 12 deletions

File tree

jumper_extension/adapters/visualizer/backends/matplotlib.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,9 @@
1212
from jumper_extension.adapters.visualizer.visualizer import PerformanceVisualizer
1313
from jumper_extension.utilities import get_available_levels
1414
from jumper_extension.logo import jumper_colors
15+
import logging
1516

17+
logger = logging.getLogger("extension")
1618

1719
def is_ipympl_backend():
1820
try:
@@ -247,6 +249,7 @@ def _draw_bali_segments(
247249
if not show_bali:
248250
return
249251
segments = self._load_bali_segments()
252+
250253
if not segments:
251254
return
252255

jumper_extension/bali_adapter.py

Lines changed: 13 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
import numpy as np
44
from itables import show
55
import logging
6+
import os
67
from jumper_extension.bali_hook import BaliResultsParser
78
from jumper_extension.core.messages import EXTENSION_INFO_MESSAGES, ExtensionInfoCode
89

@@ -45,7 +46,8 @@ def refresh_segments_from_disk(self, pid: int) -> int:
4546
Refresh BALI segments from disk for the given process ID.
4647
"""
4748
segments = self.parser.collect_all_bali_segments(pid)
48-
49+
with open("segments.txt", "w") as file:
50+
file.write(str(segments))
4951
# Build DataFrame directly from segments and align to canonical column order
5052
df = pd.DataFrame(segments)
5153
self._segments_df = df.reindex(columns=self._segments_df.columns)
@@ -117,7 +119,7 @@ def _trapz(values):
117119
powers = np.asarray(values["gpu_power_avg"], dtype=float)
118120
if len(times) < 2:
119121
return 0.0
120-
return float(np.trapz(powers, times))
122+
return float(np.trapezoid(powers, times))
121123

122124
def _safe_div(a, b):
123125
return a / b if b else None
@@ -284,7 +286,7 @@ class BaliVisualizationMixin:
284286
def __init__(self, *args, bali_adapter=None, **kwargs):
285287
super().__init__(*args, **kwargs)
286288
self.bali_adapter = bali_adapter or BaliAdapter()
287-
self._compressed_bali_segments = []
289+
self._compressed_bali_segments = None
288290
self._cached_bali_segments = None
289291

290292
def _load_bali_segments(self) -> List[Dict]:
@@ -296,11 +298,16 @@ def _load_bali_segments(self) -> List[Dict]:
296298
bali_pid = getattr(self.monitor, "bali_pid_directory", None) \
297299
or getattr(self.monitor, "pid", None)
298300
if bali_pid is None:
299-
self._cached_bali_segments = []
300-
return self._cached_bali_segments
301+
bali_pid = os.getpid()
302+
#self._cached_bali_segments = None
303+
#return self._cached_bali_segments
304+
305+
logger.info(f"BALI PID used: {bali_pid}")
306+
301307
self._cached_bali_segments = self.bali_adapter.get_segments_for_visualization(
302308
bali_pid)
303-
logging.info(f"cached segments: {self._cached_bali_segments}")
309+
logger.info(f"cached segments: {self._cached_bali_segments}")
310+
304311
return self._cached_bali_segments
305312

306313
def _invalidate_bali_cache(self):

jumper_extension/bali_hook.py

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@ def __init__(self, base_search_path: str = "."):
1616
)
1717
self.colormap_energy = mpl.colors.LinearSegmentedColormap.from_list(
1818
"yellow_to_red",
19-
['#EADFB4','#F6995C','#874C62']
19+
['#EADFB4','#874C62']
2020
)
2121

2222
def _find_bali_directories(self, pid: int) -> List[str]:
@@ -52,8 +52,8 @@ def extract_segment(
5252
framework_data = benchmark_data[framework]
5353

5454
for iteration_key, iteration_data in framework_data.items():
55-
start_time = iteration_data.get("start_time")
56-
end_time = iteration_data.get("end_time")
55+
start_time = iteration_data.get("start_timestamp")
56+
end_time = iteration_data.get("end_timestamp")
5757
# ``generation_time`` is the duration of the text-generation
5858
# phase; ``tokenize_time`` and ``setup_time`` are also
5959
# durations (not absolute timestamps).
@@ -114,8 +114,8 @@ def extract_error_segments(self, error_data: Dict, config_data: Dict) -> List[Di
114114
segments = []
115115
if error_data:
116116
for framework, error_info in error_data.items():
117-
start_time = error_info.get("start_time")
118-
end_time = error_info.get("end_time")
117+
start_time = error_info.get("start_timestamp")
118+
end_time = error_info.get("end_timestamp")
119119

120120
segments.append({
121121
"start_time": start_time,
@@ -202,4 +202,4 @@ def get_color_for_energy_efficiency(
202202
normalized = max(
203203
0.0, min(1.0, (tokens_per_sec - vmin) / (vmax - vmin))
204204
)
205-
return self.colormap_energy(normalized)
205+
return self.colormap_energy(normalized)

0 commit comments

Comments
 (0)