| 83 | |
| 84 | |
| 85 | class GenerationRequest: |
| 86 | |
| 87 | def __init__( |
| 88 | self, |
| 89 | prompt_token_ids: Union[torch.Tensor, np.ndarray, |
| 90 | Union[List[int], List[List[int]]]], |
| 91 | sampling_params: SamplingParams, |
| 92 | query_token_ids: Optional[Union[torch.Tensor, np.ndarray, list]] = None, |
| 93 | lora_request: Optional[LoRARequest] = None, |
| 94 | prompt_adapter_request: Optional[PromptAdapterRequest] = None, |
| 95 | streaming: bool = False, |
| 96 | kv_cache_retention_config: Optional[KvCacheRetentionConfig] = None, |
| 97 | disaggregated_params: Optional[DisaggregatedParams] = None, |
| 98 | trace_headers: Optional[Mapping[str, str]] = None, |
| 99 | postproc_params: Optional[PostprocParams] = None, |
| 100 | multimodal_params: Optional[MultimodalParams] = None, |
| 101 | scheduling_params: Optional[SchedulingParams] = None, |
| 102 | cache_salt_id: Optional[int] = None, |
| 103 | arrival_time: Optional[float] = None, |
| 104 | ): |
| 105 | if isinstance(prompt_token_ids, list): |
| 106 | self.prompt_token_ids = prompt_token_ids |
| 107 | self.query_token_ids = query_token_ids |
| 108 | elif isinstance(prompt_token_ids, (torch.Tensor, np.ndarray)): |
| 109 | self.prompt_token_ids = prompt_token_ids.tolist() |
| 110 | if query_token_ids: |
| 111 | self.query_token_ids = query_token_ids.tolist() |
| 112 | else: |
| 113 | raise TypeError( |
| 114 | f"prompt_token_ids ({prompt_token_ids}) should be an instance of torch.Tensor, np.ndarray or list" |
| 115 | ) |
| 116 | |
| 117 | # NOTE: Exercise caution when adding memory intense attributes, because the current implementation might lead to leaks without manual cleanup. |
| 118 | # Refer to https://github.com/NVIDIA/TensorRT-LLM/pull/5029#discussion_r2141859873 for details. |
| 119 | self.sampling_params = sampling_params |
| 120 | self.postproc_params = postproc_params |
| 121 | self.lora_request = lora_request |
| 122 | self.prompt_adapter_request = prompt_adapter_request |
| 123 | self.streaming = streaming |
| 124 | self.multimodal_params = multimodal_params |
| 125 | self.kv_cache_retention_config = kv_cache_retention_config |
| 126 | self.id: Optional[int] = None |
| 127 | self.disaggregated_params = disaggregated_params |
| 128 | self.trace_headers = trace_headers |
| 129 | self.scheduling_params = scheduling_params |
| 130 | self.cache_salt_id = cache_salt_id |
| 131 | self.arrival_time = arrival_time |
| 132 | |
| 133 | def set_id(self, id): |
| 134 | assert self.id is None, f"Request ID is already set: {self.id}" |
| 135 | self.id = id |
| 136 | return self |
| 137 | |
| 138 | |
| 139 | class CancellingRequest: |
no outgoing calls