diff --git a/src/openjarvis/telemetry/instrumented_engine.py b/src/openjarvis/telemetry/instrumented_engine.py index bc01a03c..ea0d5a00 100644 --- a/src/openjarvis/telemetry/instrumented_engine.py +++ b/src/openjarvis/telemetry/instrumented_engine.py @@ -173,6 +173,12 @@ class InstrumentedEngine(InferenceEngine): if completion_tokens > 0 and decode_latency > 0 else 0.0 ) + # --- Tier 4: Per-inference efficiency --- + tokens_per_joule = ( + completion_tokens / energy_joules + if energy_joules > 0 and completion_tokens > 0 else 0.0 + ) + engine_id = getattr(self._inner, "engine_id", "unknown") record = TelemetryRecord( @@ -201,6 +207,7 @@ class InstrumentedEngine(InferenceEngine): cpu_energy_joules=cpu_energy_joules, gpu_energy_joules=gpu_energy_joules, dram_energy_joules=dram_energy_joules, + tokens_per_joule=tokens_per_joule, ) event_data = { @@ -364,6 +371,12 @@ class InstrumentedEngine(InferenceEngine): prefill_energy = energy_joules * prefill_frac decode_energy = energy_joules * (1.0 - prefill_frac) + # Per-inference efficiency + tokens_per_joule = ( + token_count / energy_joules + if energy_joules > 0 and token_count > 0 else 0.0 + ) + engine_id = getattr(self._inner, "engine_id", "unknown") record = TelemetryRecord( @@ -397,6 +410,7 @@ class InstrumentedEngine(InferenceEngine): cpu_energy_joules=cpu_energy_joules, gpu_energy_joules=gpu_energy_joules, dram_energy_joules=dram_energy_joules, + tokens_per_joule=tokens_per_joule, ) event_data = { diff --git a/tests/telemetry/test_instrumented_engine.py b/tests/telemetry/test_instrumented_engine.py index b708ff90..2795ad9e 100644 --- a/tests/telemetry/test_instrumented_engine.py +++ b/tests/telemetry/test_instrumented_engine.py @@ -120,3 +120,45 @@ class TestInstrumentedEngine: def test_engine_id_attribute(self, mock_engine, bus): ie = InstrumentedEngine(mock_engine, bus) assert ie.engine_id == "instrumented" + + +class TestTokensPerJoule: + def test_tokens_per_joule_zero_without_energy(self, mock_engine, bus): + """tokens_per_joule is 0.0 when no energy monitor is available.""" + ie = InstrumentedEngine(mock_engine, bus) + messages = [Message(role=Role.USER, content="Hi")] + ie.generate(messages, model="test") + + tel_events = [ + e for e in bus.history + if e.event_type == EventType.TELEMETRY_RECORD + ] + record = tel_events[0].data["record"] + assert record.tokens_per_joule == 0.0 + + def test_tokens_per_joule_formula_via_record(self): + """Verify the formula: tokens_per_joule = completion_tokens / energy_joules.""" + from openjarvis.core.types import TelemetryRecord + + # Direct construction — verifies the field accepts computed values + rec = TelemetryRecord( + timestamp=1.0, + model_id="test", + completion_tokens=50, + energy_joules=2.5, + tokens_per_joule=50.0 / 2.5, # = 20.0 + ) + assert rec.tokens_per_joule == pytest.approx(20.0) + + def test_tokens_per_joule_zero_when_no_tokens(self): + """tokens_per_joule is 0.0 when completion_tokens is 0.""" + from openjarvis.core.types import TelemetryRecord + + rec = TelemetryRecord( + timestamp=1.0, + model_id="test", + completion_tokens=0, + energy_joules=5.0, + tokens_per_joule=0.0, + ) + assert rec.tokens_per_joule == 0.0