Generate tokens with post-hoc early exit evaluation. Runs the full model forward with output_hidden_states=True, then evaluates routers post-hoc to decide which layer's hidden state to use for logits. All layers run every step (preserving KV cache correctness), but the outpu
(
self,
input_ids: torch.Tensor,
max_new_tokens: int,
temperature: float,
top_k: int,
top_p: float,
)
| 221 | ) |
| 222 | |
| 223 | def _generate_with_skipping( |
| 224 | self, |
| 225 | input_ids: torch.Tensor, |
| 226 | max_new_tokens: int, |
| 227 | temperature: float, |
| 228 | top_k: int, |
| 229 | top_p: float, |
| 230 | ) -> torch.Tensor: |
| 231 | """Generate tokens with post-hoc early exit evaluation. |
| 232 | |
| 233 | Runs the full model forward with output_hidden_states=True, then evaluates |
| 234 | routers post-hoc to decide which layer's hidden state to use for logits. |
| 235 | All layers run every step (preserving KV cache correctness), but the output |
| 236 | comes from the earliest layer whose router fires. |
| 237 | |
| 238 | Compatible with all transformers versions (no hooks or exceptions). |
| 239 | """ |
| 240 | device = input_ids.device |
| 241 | B = input_ids.shape[0] |
| 242 | n_layers = len(self._layers) |
| 243 | generated = input_ids.clone() |
| 244 | gen_stats = ExitStats(total_tokens=0) |
| 245 | |
| 246 | # Prefill: run full model |
| 247 | prefill_out = self.model( |
| 248 | generated, use_cache=True, return_dict=True, |
| 249 | ) |
| 250 | past_key_values = prefill_out.past_key_values |
| 251 | next_logits = prefill_out.logits[:, -1, :] |
| 252 | |
| 253 | for step in range(max_new_tokens): |
| 254 | next_token = self._sample_next_token(next_logits, temperature, top_k, top_p) |
| 255 | generated = torch.cat([generated, next_token], dim=-1) |
| 256 | |
| 257 | if (next_token == 2).all(): |
| 258 | break |
| 259 | |
| 260 | gen_stats.total_tokens += B |
| 261 | |
| 262 | # Run full forward with hidden states for post-hoc router evaluation |
| 263 | out = self.model( |
| 264 | next_token, |
| 265 | past_key_values=past_key_values, |
| 266 | use_cache=True, |
| 267 | return_dict=True, |
| 268 | output_hidden_states=True, |
| 269 | ) |
| 270 | past_key_values = out.past_key_values |
| 271 | all_hidden = out.hidden_states # [emb, layer_0, ..., layer_N] |
| 272 | |
| 273 | # Evaluate routers post-hoc to find earliest exit |
| 274 | exit_hidden = None |
| 275 | exit_layer = None |
| 276 | for layer_idx in sorted(self.routers.keys()): |
| 277 | if layer_idx < self.config.min_layers: |
| 278 | continue |
| 279 | hidden = all_hidden[layer_idx + 1] # +1 for embedding offset |
| 280 | scores = self._score_router(hidden, layer_idx) |
no test coverage detected