MCPcopy Create free account
hub / github.com/NVIDIA/TensorRT-LLM / GenerationRequest

Class GenerationRequest

tensorrt_llm/executor/request.py:85–136  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

83
84
85class 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
139class CancellingRequest:

Calls

no outgoing calls