ore details on performance profiling with nsys). import nvtx from torch.cuda import profiler @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): if i == 30: # start nsys profiler torch.cuda.synchronize() profiler.start() elif i == 50: # stop nsys profiler torch.cuda.synchronize() profiler.stop() with nvtx.annotate(f"Step {i+1}", color="blue"): with nvtx.annotate("Model Forward", color="green"): 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 ) past_key_values = outputs.past_key_values logits = outputs.logits[:, -1, :] new_tokens = torch.argmax(logits, dim=-1) all_tokens = torch.cat( [all_tokens, new_tokens.unsqueeze(-1)], dim=-1 ) finished |= (new_tokens == config.eos_token_id) stop_gpu = torch.all(finished) with nvtx.annotate("Check Stop Condition", color="red"): # checking stop condition if stop_gpu.item(): print(f"All sequences finished at step {i+1}") break return all_tokens