@torch.inference_mode() def generate_sequence_pipelined( 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) past_key_values = None # 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 ) # Initialize cache position tracking for static cache cache_positions = torch.arange(max_seqlen, device=DEVICE) # Dual streams for pipelining streams = [torch.cuda.Stream(), torch.cuda.Stream()] stop_host = [ torch.tensor(False, pin_memory=True), torch.tensor(False, pin_memory=True) ] for i in range(max_seqlen): curr_idx, prev_idx = i % 2, (i+1) % 2 curr_s, prev_s = streams[curr_idx], streams[prev_idx] # Launch iteration i in current stream with torch.cuda.stream(curr_s): # program stream to wait for previous stream to complete curr_s.wait_stream(prev_s) 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) stop_host[curr_idx].copy_(stop_gpu, non_blocking=True) # Check previous iteration's stop signal torch.cuda.current_stream().wait_stream(prev_s) if stop_host[prev_idx].item(): print(f"All sequences finished at step {i}") break return all_tokens