Unify forward mode (#1360)
This commit is contained in:
@@ -103,7 +103,7 @@ class LogitsProcessor(nn.Module):
|
||||
|
||||
@staticmethod
|
||||
def get_top_logprobs(all_logprobs: torch.Tensor, logits_metadata: LogitsMetadata):
|
||||
if logits_metadata.forward_mode == ForwardMode.DECODE:
|
||||
if logits_metadata.forward_mode.is_decode():
|
||||
output_top_logprobs = []
|
||||
max_k = max(logits_metadata.top_logprobs_nums)
|
||||
ret = all_logprobs.topk(max_k, dim=1)
|
||||
@@ -163,7 +163,7 @@ class LogitsProcessor(nn.Module):
|
||||
assert isinstance(logits_metadata, LogitsMetadata)
|
||||
|
||||
# Get the last hidden states and last logits for the next token prediction
|
||||
if logits_metadata.forward_mode == ForwardMode.DECODE:
|
||||
if logits_metadata.forward_mode.is_decode():
|
||||
last_index = None
|
||||
last_hidden = hidden_states
|
||||
else:
|
||||
@@ -195,7 +195,7 @@ class LogitsProcessor(nn.Module):
|
||||
)
|
||||
else:
|
||||
# When logprob is requested, compute the logits for all tokens.
|
||||
if logits_metadata.forward_mode == ForwardMode.DECODE:
|
||||
if logits_metadata.forward_mode.is_decode():
|
||||
last_logprobs = torch.nn.functional.log_softmax(last_logits, dim=-1)
|
||||
|
||||
# Get the logprob of top-k tokens
|
||||
|
||||
@@ -197,9 +197,9 @@ class RadixAttention(nn.Module):
|
||||
k = k.view(-1, self.tp_k_head_num, self.qk_head_dim)
|
||||
v = v.view(-1, self.tp_v_head_num, self.v_head_dim)
|
||||
|
||||
if input_metadata.forward_mode == ForwardMode.EXTEND:
|
||||
if input_metadata.forward_mode.is_extend():
|
||||
return self.extend_forward(q, k, v, input_metadata)
|
||||
elif input_metadata.forward_mode == ForwardMode.DECODE:
|
||||
elif input_metadata.forward_mode.is_decode():
|
||||
return self.decode_forward(q, k, v, input_metadata)
|
||||
|
||||
def store_kv_cache(self, cache_k, cache_v, input_metadata: InputMetadata):
|
||||
|
||||
Reference in New Issue
Block a user