dantinox.profiling

The profiling module has no dependencies on training or paradigms. Both utilities can be used standalone.


FLOPs estimation

dantinox.profiling.counter.count_flops(config: ModelConfig, seq_len: int, batch_size: int = 1) → FLOPsBreakdown[source]

Estimate FLOPs for one forward pass using standard approximations.

Attention (per layer):

QKV projection : 3 × 2BTD² (two ops per multiply-add) output proj : 2BTD² attention score : 2BT²D

FFN (per layer, with optional SwiGLU gate doubling):

up : 2BT·D·(E·D) × swiglu_factor down : 2BT·(E·D)·D

Logit projection (unembed):

2BT·V·D

class dantinox.profiling.counter.FLOPsBreakdown(attention: int, ffn: int, embedding: int, total: int)[source]

Bases: object

Per-component FLOPs for a single forward pass.

attention: int
ffn: int
embedding: int
total: int

FLOPs formulas

\[ \text{Attention} = \left(4 \cdot 2BTD^2 + 2BT^2D\right) \times L \]
\[ \text{FFN} = \left(2BT \cdot D \cdot ED \cdot s_\text{swiglu} + 2BT \cdot ED \cdot D\right) \times L \]
\[ \text{Embedding} = 2BT \cdot V \cdot D \]

where \(B\) = batch, \(T\) = seq len, \(D\) = dim, \(E\) = expansion, \(L\) = layers, \(V\) = vocab, \(s_\text{swiglu} = 2\) if SwiGLU else \(1\).


Latency tracking

class dantinox.profiling.tracker.LatencyTracker(window: int = 10000)[source]

Bases: object

Legacy tracker. Prefer LatencyMetric for new code.

record(elapsed_s: float, n_tokens: int) → None[source]
measure(n_tokens: int) → Iterator[None][source]
result() → ProfilingResult[source]
reset() → None[source]
class dantinox.profiling.tracker.ProfilingResult(latency_mean_ms: float, latency_p50_ms: float, latency_p99_ms: float, throughput_tps: float, n_samples: int, total_tokens: int, flops: Any | None = None)[source]

Bases: object

Legacy result type. Prefer LatencyResult for new code.

latency_mean_ms: float
latency_p50_ms: float
latency_p99_ms: float
throughput_tps: float
n_samples: int
total_tokens: int
flops: Any | None = None
dantinox.profiling.tracker.profile_fn(fn: Callable[[...], Any], tracker: LatencyTracker, n_tokens: int) → Callable[[...], Any][source]

Legacy helper. Wraps fn to record one sample per call.


Usage example

from dantinox.profiling import LatencyTracker, count_flops, profile_fn
from dantinox.core.config import ModelConfig

# --- Analytical FLOPs (no model instance needed) ---
cfg   = ModelConfig(dim=512, n_heads=8, head_size=64, num_blocks=12, vocab_size=32_000)
flops = count_flops(cfg, seq_len=512, batch_size=4)
print(flops)
# FLOPs breakdown:
#   attention : 12.88 GFLOPs
#   ffn       : 25.77 GFLOPs
#   embedding : 0.13  GFLOPs
#   total     : 38.78 GFLOPs

# --- Wall-clock latency (JAX barrier-accurate) ---
tracker = LatencyTracker()

with tracker.measure(n_tokens=4 * 512):
    _ = model(x)

result = tracker.result()
print(f"mean: {result.latency_mean_ms:.1f} ms")
print(f"p99:  {result.latency_p99_ms:.1f} ms")
print(f"tps:  {result.throughput_tps:,.0f} tok/s")

# --- Functional wrapper ---
instrumented_generate = profile_fn(model.generate, tracker, n_tokens=256)
output = instrumented_generate(prompt, rng)

JAX synchronization

LatencyTracker.measure() calls jax.effects_barrier() before and after the measured call. This ensures all XLA-compiled operations have completed before the timer stops. Without this, JAX’s asynchronous dispatch would cause the measured time to reflect only dispatch latency, not actual computation.