refactor: clean up logits processor helpers (#34158)
This commit is contained in:
@@ -70,68 +70,6 @@ _UNQUANTIZED_LM_HEAD_METHODS = {
|
|||||||
"PackWeightMethod",
|
"PackWeightMethod",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def _has_lm_head_runtime_attrs(lm_head, attr_names: Tuple[str, ...]) -> bool:
|
|
||||||
return all(hasattr(lm_head, attr_name) for attr_name in attr_names)
|
|
||||||
|
|
||||||
|
|
||||||
def should_apply_lm_head_quant_method(lm_head, quant_method) -> bool:
|
|
||||||
if (
|
|
||||||
quant_method is None
|
|
||||||
or not hasattr(lm_head, "weight")
|
|
||||||
or not callable(getattr(quant_method, "apply", None))
|
|
||||||
):
|
|
||||||
return False
|
|
||||||
|
|
||||||
method_name = type(quant_method).__name__
|
|
||||||
if method_name in _UNQUANTIZED_LM_HEAD_METHODS:
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Some draft models share an unquantized target lm_head tensor while still
|
|
||||||
# carrying the draft model's stale ModelOpt quant_method. Only use the
|
|
||||||
# ModelOpt lm_head kernel when the runtime quantization state matches it.
|
|
||||||
if method_name == "ModelOptFp4LinearMethod":
|
|
||||||
if lm_head.weight.dtype == torch.int32 and _has_lm_head_runtime_attrs(
|
|
||||||
lm_head,
|
|
||||||
(
|
|
||||||
"weight_scale",
|
|
||||||
"weight_global_scale",
|
|
||||||
"workspace",
|
|
||||||
"input_size_per_partition",
|
|
||||||
"output_size_per_partition",
|
|
||||||
),
|
|
||||||
):
|
|
||||||
return True
|
|
||||||
return lm_head.weight.dtype == torch.uint8 and _has_lm_head_runtime_attrs(
|
|
||||||
lm_head,
|
|
||||||
(
|
|
||||||
"weight_scale_interleaved",
|
|
||||||
"alpha",
|
|
||||||
"input_scale_inv",
|
|
||||||
"input_size_per_partition",
|
|
||||||
"output_size_per_partition",
|
|
||||||
),
|
|
||||||
)
|
|
||||||
if method_name == "ModelOptNvFp4A16LinearMethod":
|
|
||||||
return lm_head.weight.dtype == torch.int32 and _has_lm_head_runtime_attrs(
|
|
||||||
lm_head,
|
|
||||||
(
|
|
||||||
"weight_scale",
|
|
||||||
"weight_global_scale",
|
|
||||||
"workspace",
|
|
||||||
"input_size_per_partition",
|
|
||||||
"output_size_per_partition",
|
|
||||||
),
|
|
||||||
)
|
|
||||||
if method_name == "ModelOptFp8LinearMethod":
|
|
||||||
return (
|
|
||||||
lm_head.weight.dtype == torch.float8_e4m3fn
|
|
||||||
and _has_lm_head_runtime_attrs(lm_head, ("weight_scale", "input_scale"))
|
|
||||||
)
|
|
||||||
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|
||||||
# When set, LogitsProcessor.forward returns an empty output and skips the
|
# When set, LogitsProcessor.forward returns an empty output and skips the
|
||||||
# LM head + tensor-parallel all-gather. FlashInfer autotune only profiles
|
# LM head + tensor-parallel all-gather. FlashInfer autotune only profiles
|
||||||
# attention/MoE/GEMM kernels, so the LM-head all-gather is wasted work --
|
# attention/MoE/GEMM kernels, so the LM-head all-gather is wasted work --
|
||||||
@@ -200,6 +138,10 @@ class LogitsProcessorOutput:
|
|||||||
## Part 5: Customized Info
|
## Part 5: Customized Info
|
||||||
customized_info: Optional[Dict[str, List[Any]]] = None
|
customized_info: Optional[Dict[str, List[Any]]] = None
|
||||||
|
|
||||||
|
## Part 6: Temporary variables
|
||||||
|
# FIXME: These fields are not logits-related but are passed through here as a
|
||||||
|
# workaround since ForwardBatch is local to forward_batch_generation().
|
||||||
|
# They should be moved to GenerationBatchResult to keep this class clean.
|
||||||
mm_input_embeds: Optional[torch.Tensor] = None
|
mm_input_embeds: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
|
||||||
@@ -460,9 +402,6 @@ class LogitsProcessor(nn.Module):
|
|||||||
return LogitsProcessorOutput(
|
return LogitsProcessorOutput(
|
||||||
next_token_logits=sampled_logits,
|
next_token_logits=sampled_logits,
|
||||||
hidden_states=hidden_states_to_store,
|
hidden_states=hidden_states_to_store,
|
||||||
# FIXME: These fields are not logits-related but are passed through here as a
|
|
||||||
# workaround since ForwardBatch is local to forward_batch_generation().
|
|
||||||
# They should be moved to GenerationBatchResult to keep this class clean.
|
|
||||||
mm_input_embeds=logits_metadata.mm_input_embeds,
|
mm_input_embeds=logits_metadata.mm_input_embeds,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -992,8 +931,66 @@ class LogitsProcessor(nn.Module):
|
|||||||
input_top_logprobs_idx=input_top_logprobs_idx,
|
input_top_logprobs_idx=input_top_logprobs_idx,
|
||||||
input_token_ids_logprobs_val=input_token_ids_logprobs_val,
|
input_token_ids_logprobs_val=input_token_ids_logprobs_val,
|
||||||
input_token_ids_logprobs_idx=input_token_ids_logprobs_idx,
|
input_token_ids_logprobs_idx=input_token_ids_logprobs_idx,
|
||||||
# FIXME: These fields are not logits-related but are passed through here as a
|
|
||||||
# workaround since ForwardBatch is local to forward_batch_generation().
|
|
||||||
# They should be moved to GenerationBatchResult to keep this class clean.
|
|
||||||
mm_input_embeds=logits_metadata.mm_input_embeds,
|
mm_input_embeds=logits_metadata.mm_input_embeds,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _has_lm_head_runtime_attrs(lm_head, attr_names: Tuple[str, ...]) -> bool:
|
||||||
|
return all(hasattr(lm_head, attr_name) for attr_name in attr_names)
|
||||||
|
|
||||||
|
|
||||||
|
def should_apply_lm_head_quant_method(lm_head, quant_method) -> bool:
|
||||||
|
if (
|
||||||
|
quant_method is None
|
||||||
|
or not hasattr(lm_head, "weight")
|
||||||
|
or not callable(getattr(quant_method, "apply", None))
|
||||||
|
):
|
||||||
|
return False
|
||||||
|
|
||||||
|
method_name = type(quant_method).__name__
|
||||||
|
if method_name in _UNQUANTIZED_LM_HEAD_METHODS:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Some draft models share an unquantized target lm_head tensor while still
|
||||||
|
# carrying the draft model's stale ModelOpt quant_method. Only use the
|
||||||
|
# ModelOpt lm_head kernel when the runtime quantization state matches it.
|
||||||
|
if method_name == "ModelOptFp4LinearMethod":
|
||||||
|
if lm_head.weight.dtype == torch.int32 and _has_lm_head_runtime_attrs(
|
||||||
|
lm_head,
|
||||||
|
(
|
||||||
|
"weight_scale",
|
||||||
|
"weight_global_scale",
|
||||||
|
"workspace",
|
||||||
|
"input_size_per_partition",
|
||||||
|
"output_size_per_partition",
|
||||||
|
),
|
||||||
|
):
|
||||||
|
return True
|
||||||
|
return lm_head.weight.dtype == torch.uint8 and _has_lm_head_runtime_attrs(
|
||||||
|
lm_head,
|
||||||
|
(
|
||||||
|
"weight_scale_interleaved",
|
||||||
|
"alpha",
|
||||||
|
"input_scale_inv",
|
||||||
|
"input_size_per_partition",
|
||||||
|
"output_size_per_partition",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
if method_name == "ModelOptNvFp4A16LinearMethod":
|
||||||
|
return lm_head.weight.dtype == torch.int32 and _has_lm_head_runtime_attrs(
|
||||||
|
lm_head,
|
||||||
|
(
|
||||||
|
"weight_scale",
|
||||||
|
"weight_global_scale",
|
||||||
|
"workspace",
|
||||||
|
"input_size_per_partition",
|
||||||
|
"output_size_per_partition",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
if method_name == "ModelOptFp8LinearMethod":
|
||||||
|
return (
|
||||||
|
lm_head.weight.dtype == torch.float8_e4m3fn
|
||||||
|
and _has_lm_head_runtime_attrs(lm_head, ("weight_scale", "input_scale"))
|
||||||
|
)
|
||||||
|
|
||||||
|
return True
|
||||||
|
|||||||
Reference in New Issue
Block a user