| 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 |
| |
| 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) |
|
|
| |
| 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 |
| ] |
|
|