File size: 13,759 Bytes
d91766b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
import os
from dataclasses import dataclass, field

from diffulex.engine.request import DllmReq
from diffulex.logger import get_logger

logger = get_logger(__name__)


def decode_token_ids_robust(tokenizer, token_ids: list[int] | None, *, skip_special_tokens: bool = False) -> str:
    if not token_ids:
        return ""
    try:
        return tokenizer.decode(token_ids, skip_special_tokens=skip_special_tokens)
    except TypeError:
        tokens = tokenizer.convert_ids_to_tokens(token_ids)
        safe = [t if t is not None else "" for t in tokens]
        return tokenizer.convert_tokens_to_string(safe)


@dataclass
class ReqStep:
    step_id: int
    step_time: float

    is_prefill: bool

    num_generated_tokens: int
    running_token_ids: list[int]

    block_size: int
    buffer_bids: list[int]
    block_trace: list[dict] = field(default_factory=list)

    def to_dict(self) -> dict:
        return dict(
            step_id=self.step_id,
            step_time=self.step_time,
            is_prefill=self.is_prefill,
            num_generated_tokens=self.num_generated_tokens,
            running_token_ids=self.running_token_ids,
            block_size=self.block_size,
            buffer_bids=self.buffer_bids,
            block_trace=self.block_trace,
        )


@dataclass
class ReqTrajectory:
    req_id: int

    token_ids: list[int]
    trajectory: list[ReqStep]

    is_truncated: bool
    max_new_tokens_reached: bool
    max_model_len_reached: bool
    max_nfe_reached: bool
    max_repetition_run_reached: bool
    eos_token_generated: bool
    completion_reason: str | None = None

    text: str = None
    # Generation-only tokens including content after EOS (when applicable); see DllmReq.full_response.
    full_token_ids: list[int] = field(default_factory=list)
    full_text: str | None = None

    def to_dict(self) -> dict:
        return dict(
            req_id=self.req_id,
            token_ids=self.token_ids,
            trajectory=[step.to_dict() for step in self.trajectory],
            is_truncated=self.is_truncated,
            max_new_tokens_reached=self.max_new_tokens_reached,
            max_model_len_reached=self.max_model_len_reached,
            max_nfe_reached=self.max_nfe_reached,
            max_repetition_run_reached=self.max_repetition_run_reached,
            eos_token_generated=self.eos_token_generated,
            completion_reason=self.completion_reason,
            text=self.text,
        )


class GenerationOutputs:
    """Accumulates generation outputs."""

    def __init__(self, num_prompts: int):
        self.trajectories: list[ReqTrajectory] = [
            ReqTrajectory(
                req_id=req_id,
                token_ids=[],
                trajectory=[],
                is_truncated=False,
                max_new_tokens_reached=False,
                max_model_len_reached=False,
                max_nfe_reached=False,
                max_repetition_run_reached=False,
                eos_token_generated=False,
                completion_reason=None,
            )
            for req_id in range(num_prompts)
        ]
        self._batch_step_count = 0
        self._batch_total_time = 0.0
        self._batch_generated_tokens = 0
        self._prefill_batch_time = 0.0
        self._prefill_batch_tokens = 0
        self._decode_batch_time = 0.0
        self._decode_batch_tokens = 0

    @property
    def batch_step_count(self) -> int:
        return self._batch_step_count

    @staticmethod
    def _mean(values: list[float]) -> float:
        return sum(values) / len(values) if values else 0.0
    
    @property
    def tpf(self) -> float:
        per_req_tpf = []
        for trajectory in self.trajectories:
            if not trajectory.trajectory:
                continue
            num_generated_tokens = sum(step.num_generated_tokens for step in trajectory.trajectory)
            per_req_tpf.append(num_generated_tokens / len(trajectory.trajectory))
        return self._mean(per_req_tpf)

    @property
    def ttft(self) -> float:
        per_req_ttft = []
        for trajectory in self.trajectories:
            elapsed = 0.0
            for step in trajectory.trajectory:
                elapsed += step.step_time
                if step.num_generated_tokens > 0:
                    per_req_ttft.append(elapsed)
                    break
        return self._mean(per_req_ttft)

    @property
    def tpot(self) -> float:
        per_req_tpot = []
        for trajectory in self.trajectories:
            total_time = sum(step.step_time for step in trajectory.trajectory)
            total_generated_tokens = sum(step.num_generated_tokens for step in trajectory.trajectory)
            if total_generated_tokens <= 1:
                continue

            elapsed = 0.0
            ttft = None
            for step in trajectory.trajectory:
                elapsed += step.step_time
                if step.num_generated_tokens > 0:
                    ttft = elapsed
                    break
            if ttft is not None:
                per_req_tpot.append((total_time - ttft) / (total_generated_tokens - 1))
        return self._mean(per_req_tpot)

    @property
    def throughput(self) -> float:
        return self._batch_generated_tokens / self._batch_total_time if self._batch_total_time > 0 else 0

    @property
    def e2e_total_time(self) -> float:
        return sum(sum(step.step_time for step in trajectory.trajectory) for trajectory in self.trajectories)

    @property
    def e2e_throughput(self) -> float:
        total_tokens = sum(len(trajectory.token_ids) for trajectory in self.trajectories)
        batch_time = max(self._batch_total_time, 0.001)
        return total_tokens / batch_time

    @property
    def prefill_throughput(self) -> float:
        return self._prefill_batch_tokens / self._prefill_batch_time if self._prefill_batch_time > 0 else 0

    @property
    def decode_throughput(self) -> float:
        return self._decode_batch_tokens / self._decode_batch_time if self._decode_batch_time > 0 else 0

    @property
    def avg_e2e_tps(self) -> float:
        """Mean of per-sample TPS (tokens/total_time), matching dInfer's np.mean(tpss)."""
        per_sample = []
        for trajectory in self.trajectories:
            total_time = sum(step.step_time for step in trajectory.trajectory)
            total_tokens = len(trajectory.token_ids)
            if total_time > 0 and total_tokens > 0:
                per_sample.append(total_tokens / total_time)
        return self._mean(per_sample)

    @property
    def avg_decode_tps(self) -> float:
        """Mean of per-sample decode TPS (decode_tokens / decode_time)."""
        per_sample = []
        for trajectory in self.trajectories:
            decode_time = sum(step.step_time for step in trajectory.trajectory if not step.is_prefill)
            decode_tokens = sum(step.num_generated_tokens for step in trajectory.trajectory if not step.is_prefill)
            if decode_time > 0 and decode_tokens > 0:
                per_sample.append(decode_tokens / decode_time)
        return self._mean(per_sample)

    @property
    def total_time(self) -> float:
        return self._batch_total_time

    def record_step(self, reqs: list[DllmReq], step_time: float, req_id_to_prompt_id: dict[int, int] | None = None):
        if reqs:
            self._batch_step_count += 1
            self._batch_total_time += step_time

            has_prefill = False
            has_decode = False
            prefill_tokens_this_step = 0
            decode_tokens_this_step = 0
            generated_tokens_this_step = 0

            for req in reqs:
                generated_tokens_this_step += req.new_tokens
                running_sequence = req.running_sequence
                if req.is_prefilling:
                    has_prefill = True
                    prefill_tokens_this_step += len(running_sequence or [])
                else:
                    has_decode = True
                    decode_tokens_this_step += req.new_tokens

            self._batch_generated_tokens += generated_tokens_this_step
            self._prefill_batch_tokens += prefill_tokens_this_step
            self._decode_batch_tokens += decode_tokens_this_step
            if has_prefill:
                self._prefill_batch_time += step_time
            if has_decode:
                self._decode_batch_time += step_time

        for req in reqs:
            prompt_idx = (req_id_to_prompt_id or {}).get(req.req_id, req.req_id)
            if prompt_idx >= len(self.trajectories):
                continue
            cur_trajectory = self.trajectories[prompt_idx]
            step_id = len(cur_trajectory.trajectory)

            # Per-block trace: mask ratio and active status
            block_trace = []
            if os.environ.get("DIFFULEX_SAVE_TRACE", "1") != "0":
                for block in req.dllm_block_buffer.dllm_blocks:
                    block_trace.append({
                        "block_id": block.block_id,
                        "is_active": block.is_active,
                        "is_dummy": block.is_dummy,
                        "num_mask_tokens": block.num_mask_tokens,
                        "mask_ratio": block.num_mask_tokens / max(block.block_size, 1),
                        "progress": block.progress,
                        "block_status": str(block.status) if hasattr(block, "status") else "?",
                    })

            cur_trajectory.trajectory.append(
                ReqStep(
                    step_id=step_id,
                    step_time=step_time,
                    is_prefill=req.is_prefilling,
                    num_generated_tokens=req.new_tokens,
                    running_token_ids=(
                        req.running_sequence.copy() if req.running_sequence is not None else []
                    ),
                    block_size=req.block_size,
                    buffer_bids=[block.block_id for block in req.dllm_block_buffer.dllm_blocks],
                    block_trace=block_trace,
                )
            )
            cur_trajectory.token_ids = req.truncated_response.copy() if req.truncated_response else []
            cur_trajectory.full_token_ids = list(req.full_response)
            cur_trajectory.is_truncated = req.is_truncated
            cur_trajectory.max_new_tokens_reached = req.max_new_tokens_reached
            cur_trajectory.max_model_len_reached = req.max_model_len_reached
            cur_trajectory.max_nfe_reached = req.max_nfe_reached
            cur_trajectory.max_repetition_run_reached = req.max_repetition_run_reached
            cur_trajectory.eos_token_generated = req.eos_token_generated
            cur_trajectory.completion_reason = req.completion_reason

    def postfix(self) -> dict:
        return dict(
            tpf=f"{self.tpf:.2f}tok/step",
            ttft=f"{self.ttft:.2f}s",
            tpot=f"{self.tpot:.2f}s",
            e2eps=f"{self.e2e_throughput:.2f}tok/s",
            ptps=f"{self.prefill_throughput:.2f}tok/s",
            dtps=f"{self.decode_throughput:.2f}tok/s",
        )

    def fast_postfix(self) -> dict:
        """Lightweight postfix using pre-accumulated counters — O(1) per call."""
        steps = max(self._batch_step_count, 1)
        elapsed = max(self._batch_total_time, 0.001)
        decode_elapsed = max(self._decode_batch_time, 0.001)
        return dict(
            tpf=f"{self._batch_generated_tokens / steps:.2f}tok/step",
            dtps=f"{self._decode_batch_tokens / decode_elapsed:.2f}tok/s",
            e2eps=f"{self._batch_generated_tokens / elapsed:.2f}tok/s",
        )

    def log_summary(self):
        logger.info("--------------------------------")
        logger.info("Generation Outputs Summary:")
        logger.info("--------------------------------")
        logger.info(f"Total Tokens: {sum(len(trajectory.token_ids) for trajectory in self.trajectories)} toks")
        logger.info(f"Total NFEs: {self.batch_step_count} nfes (steps)")
        logger.info(f"Total Time: {self.total_time} sec")
        logger.info(f"E2E Time: {self.e2e_total_time} sec")
        logger.info(f"TPF: {self.tpf:.2f} tok/step")
        logger.info(f"TTFT: {self.ttft:.2f} sec")
        logger.info(f"TPOT: {self.tpot:.2f} sec")
        logger.info(f"E2E Throughput: {self.e2e_throughput:.2f} tok/sec")
        logger.info(f"Prefill Throughput: {self.prefill_throughput:.2f} tok/sec")
        logger.info(f"Decode Throughput: {self.decode_throughput:.2f} tok/sec")
        logger.info(f"Avg E2E TPS (per-sample mean): {self.avg_e2e_tps:.2f} tok/sec")
        logger.info(f"Avg Decode TPS (per-sample mean): {self.avg_decode_tps:.2f} tok/sec")
        logger.info("--------------------------------")

    def convert_to_text(self, tokenizer):
        eos = getattr(tokenizer, "eos_token", None) or ""
        for trajectory in self.trajectories:
            gen_full = trajectory.full_token_ids if trajectory.full_token_ids else trajectory.token_ids
            raw_full = decode_token_ids_robust(tokenizer, gen_full)
            trajectory.full_text = raw_full
            raw_trunc = decode_token_ids_robust(tokenizer, trajectory.token_ids)
            trajectory.text = raw_trunc.split(eos)[0] if eos else raw_trunc

    def to_benchmark_format(self) -> list[dict]:
        """Convert to list of dicts expected by diffulex_bench: text, token_ids, nfe."""
        return [
            dict(
                text=t.text or "",
                full_text=(t.full_text if t.full_text is not None else t.text or ""),
                token_ids=t.token_ids if t.token_ids is not None else [],
                nfe=len(t.trajectory),
            )
            for t in self.trajectories
        ]