Tracks per-layer exit statistics.
| 11 | |
| 12 | @dataclass |
| 13 | class ExitStats: |
| 14 | """Tracks per-layer exit statistics.""" |
| 15 | total_tokens: int = 0 |
| 16 | exits_per_layer: Dict[int, int] = field(default_factory=dict) |
| 17 | remaining_tokens: int = 0 |
| 18 | |
| 19 | @property |
| 20 | def total_exited(self) -> int: |
| 21 | return sum(self.exits_per_layer.values()) |
| 22 | |
| 23 | @property |
| 24 | def exit_rate(self) -> float: |
| 25 | if self.total_tokens == 0: |
| 26 | return 0.0 |
| 27 | return self.total_exited / self.total_tokens |
| 28 | |
| 29 | def summary(self) -> str: |
| 30 | lines = [f"Total tokens: {self.total_tokens}, Exited: {self.total_exited} ({self.exit_rate:.1%})"] |
| 31 | for layer_idx in sorted(self.exits_per_layer): |
| 32 | count = self.exits_per_layer[layer_idx] |
| 33 | pct = count / self.total_tokens * 100 if self.total_tokens else 0 |
| 34 | lines.append(f" Layer {layer_idx}: {count} exits ({pct:.1f}%)") |
| 35 | lines.append(f" Ran all layers: {self.remaining_tokens}") |
| 36 | return "\n".join(lines) |
| 37 | |
| 38 | |
| 39 | class SkipScheduler: |
no outgoing calls
no test coverage detected