MCPcopy Create free account
hub / github.com/RightNow-AI/TIDE / _generate_with_skipping

Method _generate_with_skipping

python/TIDE/runtime.py:223–313  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

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)

Callers 1

generateMethod · 0.95

Calls 4

_sample_next_tokenMethod · 0.95
_score_routerMethod · 0.95
ExitStatsClass · 0.90
toMethod · 0.80

Tested by

no test coverage detected