Skip to content

Commit c1747d5

Browse files
author
Elias Werner
committed
Merge branch 'bali-hook' of https://github.com/ScaDS/jumper_jupyter_performance into bali-hook
# Conflicts: # jumper_extension/bali_adapter.py # jumper_extension/bali_hook.py
2 parents 79e902f + 5818b06 commit c1747d5

2 files changed

Lines changed: 15 additions & 6 deletions

File tree

jumper_extension/adapters/visualizer/backends/matplotlib.py

Lines changed: 4 additions & 1 deletion
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:
@@ -263,6 +265,7 @@ def _draw_bali_segments(
263265
if not show_bali:
264266
return
265267
segments = self._load_bali_segments()
268+
266269
if not segments:
267270
return
268271

@@ -358,7 +361,7 @@ def _draw_bali_segments(
358361
"Input Length": s.get("input_len", "n/a"),
359362
"Output Length": s.get("output_len", "n/a"),
360363
"Output Tokens per Second": (
361-
f"{tps:.2f}" if tps else "NaN"
364+
f"{s.get('tokens_per_sec'):.2f}" if tps else "NaN"
362365
),
363366
"Segment Throughput (Tok/s)": s.get(
364367
"segment_throughput", "n/a"

jumper_extension/bali_adapter.py

Lines changed: 11 additions & 5 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,8 +298,12 @@ 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)
303309
logging.debug("cached segments: %s", self._cached_bali_segments)

0 commit comments

Comments
 (0)