che: from transformers import StaticCache @torch.inference_mode() def generate_sequence( model, max_seqlen, batch_size, use_cache=False, use_static_cache=False ): # Initialize prompts with BOS token all_tokens = torch.full( (batch_size, 1), config.bos_token_id, device=DEVICE, dtype=torch.long ) finished = torch.zeros(batch_size, device=DEVICE, dtype=torch.bool) # Initialize static cache if requested if use_cache and use_static_cache: past_key_values = StaticCache( config=config, max_batch_size=batch_size, max_cache_len=max_seqlen, device=DEVICE, dtype=model.dtype ) else: past_key_values = None # Initialize cache position tracking for static cache cache_positions = torch.arange(max_seqlen, device=DEVICE) for i in range(max_seqlen): current_input = ( all_tokens if past_key_values is None else all_tokens[:, -1:] ) cache_position = ( cache_positions[i:i+1] if use_static_cache else None ) outputs = model( current_input, past_key_values=past_key_values, cache_position=cache_position, use_cache=use_cache ) # update cache for next step past_key_values = outputs.past_key_values logits = outputs.logits[:, -1, :] new_tokens = torch.argmax(logits, dim=-1) # append new token to sequence all_tokens = torch.cat( [all_tokens, new_tokens.unsqueeze(-1)], dim=-1 ) finished |= (new_tokens == config.eos_token_id) stop_gpu = torch.all(finished) # checking stop condition if stop_gpu.item(): print(f"All sequences finished at step {i+1}") break return all_tokens