Enhance comments in set_extend_input_len method (#16130)
This commit is contained in:
@@ -1114,10 +1114,17 @@ class Req:
|
|||||||
self.has_log_time_stats = True
|
self.has_log_time_stats = True
|
||||||
|
|
||||||
def set_extend_input_len(self, extend_input_len: int):
|
def set_extend_input_len(self, extend_input_len: int):
|
||||||
|
# Setting extend_input_len and computing the relative logprob_start_len in an extend batch
|
||||||
|
#
|
||||||
|
# Key variables:
|
||||||
|
# - logprob_start_len: Absolute position in full sequence where logprob computation begins
|
||||||
|
# - extend_logprob_start_len: Relative position within current extend batch where logprob computation begins
|
||||||
|
# - extend_input_len: Number of tokens that need to be processed in this extend batch
|
||||||
self.extend_input_len = extend_input_len
|
self.extend_input_len = extend_input_len
|
||||||
if self.logprob_start_len == -1:
|
if self.logprob_start_len == -1:
|
||||||
logprob_start_len = len(self.fill_ids) - 1
|
logprob_start_len = len(self.fill_ids) - 1
|
||||||
else:
|
else:
|
||||||
|
# logprob_start_len should be at least the length of the prefix indices
|
||||||
logprob_start_len = max(self.logprob_start_len, len(self.prefix_indices))
|
logprob_start_len = max(self.logprob_start_len, len(self.prefix_indices))
|
||||||
self.extend_logprob_start_len = min(
|
self.extend_logprob_start_len = min(
|
||||||
logprob_start_len - len(self.prefix_indices),
|
logprob_start_len - len(self.prefix_indices),
|
||||||
|
|||||||
Reference in New Issue
Block a user