File size: 1,795 Bytes
a066584
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
# Copyright 2026 Modilify
# SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
"""Optional per-denoise phase timers. Enabled only with --profile."""

from __future__ import annotations

from collections import defaultdict
import time

import mlx.core as mx

PHASES = (
    "latent",
    "attn",
    "moe",
    "lm_head",
    "softmax",
    "commit",
    "sync",
    "update_cache",
)


class DenoiseProfiler:
    def __init__(self) -> None:
        self.totals = defaultdict(float)
        self.steps = 0

    def add(self, phase: str, seconds: float) -> None:
        self.totals[phase] += float(seconds)

    def finish_step(self) -> None:
        self.steps += 1

    def measure(self, phase: str, *arrays: mx.array):
        mx.eval(*arrays)
        started = time.perf_counter()

        class _Span:
            def __init__(self, profiler: DenoiseProfiler, name: str) -> None:
                self.profiler = profiler
                self.name = name
                self.started = started

            def done(self, *outputs: mx.array) -> None:
                if outputs:
                    mx.eval(*outputs)
                self.profiler.add(self.name, time.perf_counter() - self.started)

        return _Span(self, phase)

    def summary(self) -> str:
        counted = sum(self.totals[name] for name in PHASES)
        lines = [
            f"[profile] steps={self.steps} accounted={counted:.3f}s",
        ]
        for name in PHASES:
            value = self.totals[name]
            share = (100.0 * value / counted) if counted else 0.0
            per = (value / self.steps) if self.steps else 0.0
            lines.append(
                f"[profile] {name:12s}  {value:7.3f}s  {share:5.1f}%  {per*1000:6.1f} ms/step"
            )
        return "\n".join(lines)