diff --git a/cpp/tensorrt_llm/nanobind/thop/bindings.cpp b/cpp/tensorrt_llm/nanobind/thop/bindings.cpp index 6577d1cf185d..5f089c462bf4 100644 --- a/cpp/tensorrt_llm/nanobind/thop/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/thop/bindings.cpp @@ -39,9 +39,9 @@ void initBindings(nb::module_& m) m.def("attention", &torch_ext::attention, // Parameters with default values using std::nullopt for optional arguments nb::arg("q"), nb::arg("k") = std::nullopt, nb::arg("v") = std::nullopt, nb::arg("output"), - nb::arg("output_sf") = std::nullopt, nb::arg("out_dtype") = std::nullopt, nb::arg("workspace_") = std::nullopt, - nb::arg("sequence_length"), nb::arg("host_past_key_value_lengths"), nb::arg("host_total_kv_lens"), - nb::arg("context_lengths"), nb::arg("host_context_lengths"), nb::arg("host_request_types"), + nb::arg("output_sf") = std::nullopt, nb::arg("workspace_") = std::nullopt, nb::arg("sequence_length"), + nb::arg("host_past_key_value_lengths"), nb::arg("host_total_kv_lens"), nb::arg("context_lengths"), + nb::arg("host_context_lengths"), nb::arg("host_request_types"), nb::arg("kv_cache_block_offsets") = std::nullopt, nb::arg("host_kv_cache_block_offsets") = std::nullopt, nb::arg("host_kv_cache_pool_pointers") = std::nullopt, nb::arg("host_kv_cache_pool_mapping") = std::nullopt, nb::arg("cache_indirection") = std::nullopt, nb::arg("kv_scale_orig_quant") = std::nullopt, diff --git a/cpp/tensorrt_llm/pybind/thop/bindings.cpp b/cpp/tensorrt_llm/pybind/thop/bindings.cpp index 8cdc8a99829b..f1469927ce5e 100644 --- a/cpp/tensorrt_llm/pybind/thop/bindings.cpp +++ b/cpp/tensorrt_llm/pybind/thop/bindings.cpp @@ -39,9 +39,9 @@ void initBindings(pybind11::module_& m) m.def("attention", &torch_ext::attention, // Parameters with default values using std::nullopt for optional arguments py::arg("q"), py::arg("k") = std::nullopt, py::arg("v") = std::nullopt, py::arg("output"), - py::arg("output_sf") = std::nullopt, py::arg("out_dtype") = std::nullopt, py::arg("workspace_") = std::nullopt, - py::arg("sequence_length"), py::arg("host_past_key_value_lengths"), py::arg("host_total_kv_lens"), - py::arg("context_lengths"), py::arg("host_context_lengths"), py::arg("host_request_types"), + py::arg("output_sf") = std::nullopt, py::arg("workspace_") = std::nullopt, py::arg("sequence_length"), + py::arg("host_past_key_value_lengths"), py::arg("host_total_kv_lens"), py::arg("context_lengths"), + py::arg("host_context_lengths"), py::arg("host_request_types"), py::arg("kv_cache_block_offsets") = std::nullopt, py::arg("host_kv_cache_block_offsets") = std::nullopt, py::arg("host_kv_cache_pool_pointers") = std::nullopt, py::arg("host_kv_cache_pool_mapping") = std::nullopt, py::arg("cache_indirection") = std::nullopt, py::arg("kv_scale_orig_quant") = std::nullopt, diff --git a/cpp/tensorrt_llm/thop/attentionOp.cpp b/cpp/tensorrt_llm/thop/attentionOp.cpp index b0ee56a83f5d..688732b38393 100644 --- a/cpp/tensorrt_llm/thop/attentionOp.cpp +++ b/cpp/tensorrt_llm/thop/attentionOp.cpp @@ -603,23 +603,22 @@ using torch_ext::trtllm::attention::Runner; using torch_ext::trtllm::attention::AttentionInputType; void attention(torch::Tensor q, std::optional k, std::optional v, torch::Tensor& output, - std::optional output_sf, std::optional out_dtype, - std::optional workspace_, torch::Tensor sequence_length, torch::Tensor host_past_key_value_lengths, - torch::Tensor host_total_kv_lens, torch::Tensor context_lengths, torch::Tensor host_context_lengths, - torch::Tensor host_request_types, std::optional kv_cache_block_offsets, - std::optional host_kv_cache_block_offsets, std::optional host_kv_cache_pool_pointers, - std::optional host_kv_cache_pool_mapping, std::optional cache_indirection, - std::optional kv_scale_orig_quant, std::optional kv_scale_quant_orig, - std::optional out_scale, std::optional rotary_inv_freq, - std::optional rotary_cos_sin, std::optional latent_cache, - std::optional q_pe, std::optional block_ids_per_seq, - std::optional attention_sinks, bool const is_fused_qkv, bool const update_kv_cache, - int64_t const predicted_tokens_per_seq, int64_t const layer_idx, int64_t const num_heads, - int64_t const num_kv_heads, int64_t const head_size, std::optional const tokens_per_block, - int64_t const max_num_requests, int64_t const max_context_length, int64_t const attention_window_size, - int64_t const sink_token_length, int64_t const beam_width, int64_t const mask_type, int64_t const quant_mode, - double const q_scaling, int64_t const position_embedding_type, int64_t const rotary_embedding_dim, - double const rotary_embedding_base, int64_t const rotary_embedding_scale_type, + std::optional output_sf, std::optional workspace_, torch::Tensor sequence_length, + torch::Tensor host_past_key_value_lengths, torch::Tensor host_total_kv_lens, torch::Tensor context_lengths, + torch::Tensor host_context_lengths, torch::Tensor host_request_types, + std::optional kv_cache_block_offsets, std::optional host_kv_cache_block_offsets, + std::optional host_kv_cache_pool_pointers, std::optional host_kv_cache_pool_mapping, + std::optional cache_indirection, std::optional kv_scale_orig_quant, + std::optional kv_scale_quant_orig, std::optional out_scale, + std::optional rotary_inv_freq, std::optional rotary_cos_sin, + std::optional latent_cache, std::optional q_pe, + std::optional block_ids_per_seq, std::optional attention_sinks, + bool const is_fused_qkv, bool const update_kv_cache, int64_t const predicted_tokens_per_seq, + int64_t const layer_idx, int64_t const num_heads, int64_t const num_kv_heads, int64_t const head_size, + std::optional const tokens_per_block, int64_t const max_num_requests, int64_t const max_context_length, + int64_t const attention_window_size, int64_t const sink_token_length, int64_t const beam_width, + int64_t const mask_type, int64_t const quant_mode, double const q_scaling, int64_t const position_embedding_type, + int64_t const rotary_embedding_dim, double const rotary_embedding_base, int64_t const rotary_embedding_scale_type, std::vector rotary_embedding_scales, std::vector rotary_embedding_max_position_info, bool const use_paged_context_fmha, std::optional attention_input_type, bool is_mla_enable, std::optional chunked_prefill_buffer_batch_size, std::optional q_lora_rank, @@ -658,8 +657,10 @@ void attention(torch::Tensor q, std::optional k, std::optional k, std::optional>(); } } else if (dtype == nvinfer1::DataType::kFLOAT) { - TLLM_CHECK(!out_dtype.has_value() || out_dtype.value() == torch::kFloat32); + TLLM_CHECK(out_dtype == torch::kFloat32); runner = std::make_shared>(); } #ifdef ENABLE_BF16 @@ -696,7 +697,7 @@ void attention(torch::Tensor q, std::optional k, std::optional>(); } } diff --git a/cpp/tensorrt_llm/thop/attentionOp.h b/cpp/tensorrt_llm/thop/attentionOp.h index 8414dc067bf5..9b4751aeba64 100644 --- a/cpp/tensorrt_llm/thop/attentionOp.h +++ b/cpp/tensorrt_llm/thop/attentionOp.h @@ -39,23 +39,22 @@ namespace torch_ext * - Speculative decoding */ void attention(torch::Tensor q, std::optional k, std::optional v, torch::Tensor& output, - std::optional output_sf, std::optional out_dtype, - std::optional workspace_, torch::Tensor sequence_length, torch::Tensor host_past_key_value_lengths, - torch::Tensor host_total_kv_lens, torch::Tensor context_lengths, torch::Tensor host_context_lengths, - torch::Tensor host_request_types, std::optional kv_cache_block_offsets, - std::optional host_kv_cache_block_offsets, std::optional host_kv_cache_pool_pointers, - std::optional host_kv_cache_pool_mapping, std::optional cache_indirection, - std::optional kv_scale_orig_quant, std::optional kv_scale_quant_orig, - std::optional out_scale, std::optional rotary_inv_freq, - std::optional rotary_cos_sin, std::optional latent_cache, - std::optional q_pe, std::optional block_ids_per_seq, - std::optional attention_sinks, bool const is_fused_qkv, bool const update_kv_cache, - int64_t const predicted_tokens_per_seq, int64_t const layer_idx, int64_t const num_heads, - int64_t const num_kv_heads, int64_t const head_size, std::optional const tokens_per_block, - int64_t const max_num_requests, int64_t const max_context_length, int64_t const attention_window_size, - int64_t const sink_token_length, int64_t const beam_width, int64_t const mask_type, int64_t const quant_mode, - double const q_scaling, int64_t const position_embedding_type, int64_t const rotary_embedding_dim, - double const rotary_embedding_base, int64_t const rotary_embedding_scale_type, + std::optional output_sf, std::optional workspace_, torch::Tensor sequence_length, + torch::Tensor host_past_key_value_lengths, torch::Tensor host_total_kv_lens, torch::Tensor context_lengths, + torch::Tensor host_context_lengths, torch::Tensor host_request_types, + std::optional kv_cache_block_offsets, std::optional host_kv_cache_block_offsets, + std::optional host_kv_cache_pool_pointers, std::optional host_kv_cache_pool_mapping, + std::optional cache_indirection, std::optional kv_scale_orig_quant, + std::optional kv_scale_quant_orig, std::optional out_scale, + std::optional rotary_inv_freq, std::optional rotary_cos_sin, + std::optional latent_cache, std::optional q_pe, + std::optional block_ids_per_seq, std::optional attention_sinks, + bool const is_fused_qkv, bool const update_kv_cache, int64_t const predicted_tokens_per_seq, + int64_t const layer_idx, int64_t const num_heads, int64_t const num_kv_heads, int64_t const head_size, + std::optional const tokens_per_block, int64_t const max_num_requests, int64_t const max_context_length, + int64_t const attention_window_size, int64_t const sink_token_length, int64_t const beam_width, + int64_t const mask_type, int64_t const quant_mode, double const q_scaling, int64_t const position_embedding_type, + int64_t const rotary_embedding_dim, double const rotary_embedding_base, int64_t const rotary_embedding_scale_type, std::vector rotary_embedding_scales, std::vector rotary_embedding_max_position_info, bool const use_paged_context_fmha, std::optional attention_input_type, bool is_mla_enable, std::optional chunked_prefill_buffer_batch_size, std::optional q_lora_rank, diff --git a/tensorrt_llm/_torch/attention_backend/interface.py b/tensorrt_llm/_torch/attention_backend/interface.py index 6d3f22844077..677a441c4059 100644 --- a/tensorrt_llm/_torch/attention_backend/interface.py +++ b/tensorrt_llm/_torch/attention_backend/interface.py @@ -705,9 +705,13 @@ def support_fused_qkv(cls) -> bool: def support_mla(cls) -> bool: return False - @classmethod - def support_nvfp4_output(cls) -> bool: - return False + def create_output(self, q: torch.Tensor, **kwargs) -> List[torch.Tensor]: + """ + Create the output tensors for the attention operation. + """ + num_tokens = q.shape[0] + hidden_size = self.num_heads * self.head_dim + return [q.new_empty([num_tokens, hidden_size], dtype=q.dtype)] @dataclass(kw_only=True, unsafe_hash=True) diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index e669925a3474..9f45ca1c529d 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -2,7 +2,7 @@ import os import weakref from dataclasses import dataclass, field -from typing import TYPE_CHECKING, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, List, Optional, Tuple import torch @@ -86,6 +86,7 @@ class TrtllmAttentionWrapper: spec_bl_tree_first_sparse_mask_offset_kv: Optional[torch.Tensor] helix_position_offsets: Optional[torch.Tensor] helix_is_inactive_rank: Optional[torch.Tensor] + attention_input_type: Optional[torch.Tensor] kwargs: dict def __init__( @@ -148,6 +149,7 @@ def __init__( self.rotary_embedding_long_m_scale = rope_params.long_m_scale self.rotary_embedding_max_positions = rope_params.max_positions self.rotary_embedding_original_max_positions = rope_params.original_max_positions + self.attention_input_type = None self.kwargs = {} self.kwargs.update(kwargs) self.skip_softmax_stat = torch.zeros(2, @@ -330,18 +332,20 @@ def plan( self.skip_softmax_threshold_scale_factor_decode = skip_softmax_threshold_scale_factor_decode self.kwargs.update(kwargs) - def create_output(self, q: torch.Tensor, out_dtype: torch.dtype): + def create_output( + self, + q: torch.Tensor, + out_dtype: Optional[torch.dtype], + use_nvfp4_output: bool, + is_gen_only: bool, + ): num_tokens = q.size(0) - attention_input_type = (AttentionInputType(self.attention_input_type) - if self.attention_input_type is not None else - AttentionInputType.mixed) if out_dtype is None: out_dtype = q.dtype - is_gen_only = attention_input_type == AttentionInputType.generation_only v_head_size = self.head_size if self.is_mla_enable: v_head_size = self.kv_lora_rank if is_gen_only else self.v_head_dim - if out_dtype == torch.uint8: + if use_nvfp4_output: num_nvfp4_elements_per_container = 2 scaling_vector_size = 16 size_per_token = self.num_heads * v_head_size @@ -353,20 +357,20 @@ def create_output(self, q: torch.Tensor, out_dtype: torch.dtype): padded_row, padded_col = compute_swizzled_sf_shape( num_tokens, size_per_token // scaling_vector_size) output_sf = q.new_empty(padded_row * padded_col, dtype=torch.uint8) + return [output, output_sf] else: - output = q.new_empty((num_tokens, self.num_heads * v_head_size), - dtype=out_dtype) - output_sf = None - return output, output_sf + return [ + q.new_empty((num_tokens, self.num_heads * v_head_size), + dtype=out_dtype) + ] def run( self, q: torch.Tensor, + output: torch.Tensor, + output_sf: Optional[torch.Tensor] = None, k: Optional[torch.Tensor] = None, v: Optional[torch.Tensor] = None, - output: Optional[torch.Tensor] = None, - output_sf: Optional[torch.Tensor] = None, - out_dtype: Optional[torch.dtype] = None, is_fused_qkv: bool = True, update_kv_cache: bool = True, attention_mask: AttentionMask = PredefinedAttentionMask.CAUSAL, @@ -381,9 +385,10 @@ def run( Run the attention operation. Args: q (torch.Tensor): Query tensor with shape (num_tokens, num_heads * head_dim) or QKV tensor with shape (num_tokens, (num_heads + 2 * num_kv_heads) * head_dim). + output (torch.Tensor): Output tensor with shape. + output_sf (Optional[torch.Tensor]): Output scaling factors tensor. k (Optional[torch.Tensor]): Key tensor with shape (num_tokens, num_kv_heads * head_dim) or None if QKV tensor is provided. v (Optional[torch.Tensor]): Value tensor with shape (num_tokens, num_kv_heads * head_dim) or None if QKV tensor is provided. - out_dtype (Optional[torch.dtype]): Output data type if provided. is_fused_qkv (bool): Whether QKV tensor is provided. update_kv_cache (bool): Whether KV cache is updated. attention_mask (AttentionMask): Attention mask. See definition of AttentionMask for accepted types. Defaults to predefined causal mask. @@ -462,13 +467,6 @@ def run( else: raise ValueError("Unexpected attention mask type") - if output is None: - assert output_sf is None - output, output_sf = self.create_output(q, out_dtype) - else: - # output is provided, expect output_sf be provided as well if has NVFP4 output. - assert out_dtype is None or out_dtype != torch.uint8 or output_sf is not None - # packing parameters to avoid maxing out 64 arguments rotary_embedding_scales = [ self.rotary_embedding_scale, self.rotary_embedding_short_m_scale, @@ -505,7 +503,6 @@ def run( v, output, output_sf, - out_dtype, self.workspace, self.sequence_length, self.host_past_key_value_lengths, @@ -591,7 +588,6 @@ def run( # reset the planned states (especially tensors) to avoid memory leak self.plan() - return output, output_sf def is_nvfp4_output_kernel_available( self, @@ -1552,12 +1548,66 @@ def get_local_layer_idx(self, metadata: TrtllmAttentionMetadata) -> int: else: return metadata.kv_cache_manager.layer_offsets[self.layer_idx] + def use_nvfp4_output( + self, + metadata: TrtllmAttentionMetadata, + attention_mask: AttentionMask, + ) -> bool: + # Not running NVFP4 + if not self.has_nvfp4: + return False + + # Default enabled, but allow manual disabling through `TRTLLM_ENABLE_ATTENTION_NVFP4_OUTPUT=0` + if not os.environ.get("TRTLLM_ENABLE_ATTENTION_NVFP4_OUTPUT", + "1") == "1": + return False + + use_paged_context_fmha = ( + metadata.runtime_features.chunked_prefill + or metadata.runtime_features.cache_reuse + or metadata.runtime_features.has_speculative_draft_tokens + ) if metadata.runtime_features else False + + return self.wrapper.is_nvfp4_output_kernel_available( + tokens_per_block=metadata.tokens_per_block, + attention_mask=attention_mask, + use_paged_context_fmha=use_paged_context_fmha, + is_mla_enable=self.is_mla_enable, + ) + + def get_quantize_output_dtype( + self, use_nvfp4_output: bool) -> Optional[torch.dtype]: + if use_nvfp4_output: + # Use UINT8 as the container dtype for NVFP4. + return torch.uint8 + elif (self.has_fp8_qdq or self.has_nvfp4 or self.has_fp8_block_wise + or self.has_fp8_rowwise + or self.has_w4a8_nvfp4_fp8) and (self.has_fp8_kv_cache + or self.has_fp4_kv_cache): + return torch.float8_e4m3fn + return None + + def create_output(self, q, *, is_quantize_output: bool, + metadata: TrtllmAttentionMetadata, + attention_mask: AttentionMask, is_gen_only: bool, + **kwargs) -> List[torch.Tensor]: + use_nvfp4_output = False + out_dtype = None + if is_quantize_output: + use_nvfp4_output = self.use_nvfp4_output(metadata, attention_mask) + out_dtype = self.get_quantize_output_dtype(use_nvfp4_output) + + return self.wrapper.create_output(q, out_dtype, use_nvfp4_output, + is_gen_only) + def forward( self, q: torch.Tensor, k: Optional[torch.Tensor], v: Optional[torch.Tensor], metadata: TrtllmAttentionMetadata, + output: Optional[torch.Tensor] = None, + output_sf: Optional[torch.Tensor] = None, out_scale: Optional[torch.Tensor] = None, out_scale_sf: Optional[torch.Tensor] = None, kv_scales_sf: Optional[torch.Tensor] = None, @@ -1571,8 +1621,6 @@ def forward( attention_window_size: Optional[int] = None, softmax_stats_tensor: Optional[torch.Tensor] = None, enable_attn_nvfp4_output: bool = True, - output: Optional[torch.Tensor] = None, - output_sf: Optional[torch.Tensor] = None, attention_sinks: Optional[torch.Tensor] = None, chunked_prefill_buffer_batch_size: int = 1, cu_q_seqlens: Optional[torch.Tensor] = None, @@ -1582,7 +1630,31 @@ def forward( mla_bmm2_scale: Optional[torch.Tensor] = None, quant_q_buffer: Optional[torch.Tensor] = None, **kwargs, - ) -> Union[torch.Tensor, Tuple[torch.Tensor, Optional[torch.Tensor]]]: + ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + """ + Execute the attention operation. + Args: + q (torch.Tensor): Query tensor or QKV tensor. + k (Optional[torch.Tensor]): Key tensor or None if QKV tensor is provided. + v (Optional[torch.Tensor]): Value tensor or None if QKV tensor is provided. + metadata (TrtllmAttentionMetadata): Metadata for the attention operation. + output (Optional[torch.Tensor]): Output tensor to store the attention output. + output_sf (Optional[torch.Tensor]): Output scale factor tensor for NVFP4. + out_scale (Optional[torch.Tensor]): Scale factor tensor for quantizing output. + out_scale_sf (Optional[torch.Tensor]): Global scale factor tensor for NVFP4 for quantizingoutput. + kv_scales_sf (Optional[torch.Tensor]): KV scale factor tensor. + kv_scales_sf_inv (Optional[torch.Tensor]): KV scale factor inverse tensor. + attention_mask (AttentionMask): Attention mask. + attention_input_type (AttentionInputType): Attention input type. + latent_cache (Optional[torch.Tensor]): Latent cache tensor. + q_pe (Optional[torch.Tensor]): Q position embedding tensor. + mrope_config (Optional[dict]): Mrope configuration. + attention_window_size (Optional[int]): Attention window size. + softmax_stats_tensor (Optional[torch.Tensor]): Softmax statistics tensor. + helix_position_offsets (Optional[torch.Tensor]): Helix position offsets tensor. + attention_sinks (Optional[torch.Tensor]): Attention sinks tensor. + chunked_prefill_buffer_batch_size (int): Chunked prefill buffer batch size. + """ assert isinstance( metadata, TrtllmAttentionMetadata, @@ -1599,17 +1671,22 @@ def forward( # Context MLA uses separate qkv instead of paged_context_fmha use_paged_context_fmha = False - use_nvfp4_output = False - if enable_attn_nvfp4_output and self.has_nvfp4 and self.support_nvfp4_output( - ): - # Runtime check whether the NVFP4 output kernel is available. - use_nvfp4_output = self.wrapper.is_nvfp4_output_kernel_available( - tokens_per_block=metadata.tokens_per_block, + if output is None: + # Output is not provided. + is_gen_only = attention_input_type == AttentionInputType.generation_only + outputs = self.create_output( + q, + is_quantize_output=out_scale is not None, + metadata=metadata, attention_mask=attention_mask, use_paged_context_fmha=use_paged_context_fmha, is_mla_enable=self.is_mla_enable, + is_gen_only=is_gen_only, ) + output = outputs[0] + output_sf = outputs[1] if len(outputs) == 2 else None + sparse_kv_indices, sparse_kv_offsets, sparse_attn_indices, sparse_attn_offsets = None, None, None, None sparse_attn_indices_block_size = 1 skip_softmax_threshold_scale_factor_prefill = None @@ -1659,7 +1736,8 @@ def forward( out_scale_sf=out_scale_sf, kv_scales_sf=kv_scales_sf, kv_scales_sf_inv=kv_scales_sf_inv, - use_nvfp4_output=use_nvfp4_output, + use_nvfp4_output=output_sf + is not None, # NVFP4 output will setup output_sf tensor use_paged_context_fmha=use_paged_context_fmha, attention_input_type=attention_input_type, latent_cache=latent_cache, @@ -1695,40 +1773,27 @@ def forward( helix_position_offsets=metadata.helix_position_offsets, helix_is_inactive_rank=metadata.helix_is_inactive_rank, ) - out_dtype = None - if out_scale is not None: - if use_nvfp4_output: - # Use UINT8 as the container dtype for NVFP4. - out_dtype = torch.uint8 - elif (self.has_fp8_qdq or self.has_nvfp4 or self.has_fp8_block_wise - or self.has_fp8_rowwise - or self.has_w4a8_nvfp4_fp8) and (self.has_fp8_kv_cache - or self.has_fp4_kv_cache): - # TODO(qijun): revisit fp8_context_fmha logic - out_dtype = torch.float8_e4m3fn - - output, output_sf = self.wrapper.run( - q, - k, - v, - output=output, - output_sf=output_sf, - out_dtype=out_dtype, - is_fused_qkv=not metadata.is_cross and k is None, - update_kv_cache=not metadata.is_cross or k is not None, - attention_mask=attention_mask, - cu_q_seqlens=cu_q_seqlens, - cu_kv_seqlens=cu_kv_seqlens, - fmha_scheduler_counter=fmha_scheduler_counter, - mla_bmm1_scale=mla_bmm1_scale, - mla_bmm2_scale=mla_bmm2_scale, - quant_q_buffer=quant_q_buffer) - if use_nvfp4_output: + self.wrapper.run(q, + output, + output_sf, + k, + v, + is_fused_qkv=not metadata.is_cross and k is None, + update_kv_cache=not metadata.is_cross or k is not None, + attention_mask=attention_mask, + cu_q_seqlens=cu_q_seqlens, + cu_kv_seqlens=cu_kv_seqlens, + fmha_scheduler_counter=fmha_scheduler_counter, + mla_bmm1_scale=mla_bmm1_scale, + mla_bmm2_scale=mla_bmm2_scale, + quant_q_buffer=quant_q_buffer) + + if output_sf is None: + return output + else: return output, output_sf - return output - @classmethod def support_fused_rope(cls) -> bool: return True @@ -1741,12 +1806,6 @@ def support_fused_qkv(cls) -> bool: def support_mla(cls) -> bool: return True - @classmethod - def support_nvfp4_output(cls) -> bool: - # Default enabled, but allow manual disabling through `TRTLLM_ENABLE_ATTENTION_NVFP4_OUTPUT=0` - return os.environ.get("TRTLLM_ENABLE_ATTENTION_NVFP4_OUTPUT", - "1") == "1" - def has_cached_kv_for_mla_context( self, metadata: TrtllmAttentionMetadata, diff --git a/tensorrt_llm/_torch/compilation/utils.py b/tensorrt_llm/_torch/compilation/utils.py index 07430c60178e..dc22d0868112 100644 --- a/tensorrt_llm/_torch/compilation/utils.py +++ b/tensorrt_llm/_torch/compilation/utils.py @@ -62,6 +62,7 @@ def inplace_info(): }, torch.ops.trtllm.attn_custom_op_inplace.default: { 1: "output", + 2: "output_sf" }, torch.ops.trtllm.mla_custom_op_inplace.default: { 1: "output" diff --git a/tensorrt_llm/_torch/modules/attention.py b/tensorrt_llm/_torch/modules/attention.py index 10cce12a5d9d..47e793e48169 100644 --- a/tensorrt_llm/_torch/modules/attention.py +++ b/tensorrt_llm/_torch/modules/attention.py @@ -1,6 +1,6 @@ import math import weakref -from typing import Optional, Union, cast +from typing import List, Optional, Union, cast import torch from torch import nn @@ -15,6 +15,7 @@ FlashInferAttentionMetadata, TrtllmAttention, TrtllmAttentionMetadata) from ..attention_backend.interface import (AttentionBackend, AttentionMask, + CustomAttentionMask, PositionalEmbeddingParams, PredefinedAttentionMask) from ..attention_backend.sparse.dsa import ( @@ -87,8 +88,25 @@ def maybe_compiled_cat(tensors, dim): return torch.cat(tensors, dim) +def create_attn_outputs_impl(q: torch.Tensor, attention_mask: str, + layer_idx: str) -> List[torch.Tensor]: + metadata, attn_layer = extract_extra_attrs(layer_idx, "attn") + return attn_layer.create_output(q, metadata, attention_mask) + + +@torch.library.custom_op("trtllm::create_attn_outputs", mutates_args=()) +def create_attn_outputs(q: torch.Tensor, attention_mask: str, + layer_idx: str) -> List[torch.Tensor]: + return create_attn_outputs_impl(q, attention_mask, layer_idx) + + +@create_attn_outputs.register_fake +def _(q, attention_mask, layer_idx): + return create_attn_outputs_impl(q, attention_mask, layer_idx) + + @torch.library.custom_op("trtllm::attn_custom_op_inplace", - mutates_args=("output", )) + mutates_args=("output", "output_sf")) def attn_custom_op_inplace( q: torch.Tensor, k: Optional[torch.Tensor], @@ -101,20 +119,25 @@ def attn_custom_op_inplace( attention_sinks: Optional[torch.Tensor], layer_idx: str, output: torch.Tensor, + output_sf: Optional[torch.Tensor], ) -> None: metadata, attn_layer = extract_extra_attrs(layer_idx, "attn") + mask = PredefinedAttentionMask( + attention_mask + ) if attention_mask != CustomAttentionMask.CUSTOM else CustomAttentionMask( + attention_mask) # NVFP4 output cannot be supported by torch compile for TRTLLM backend. attn_layer._attn_impl(q, k, v, metadata, - PredefinedAttentionMask(attention_mask), + mask, mrope_rotary_cos_sin, mrope_position_deltas, attention_window_size, attention_mask_data, - enable_attn_nvfp4_output=False, output=output, + output_sf=output_sf, attention_sinks=attention_sinks) @@ -170,6 +193,11 @@ def __init__( if config is not None: if "attn_layers" not in config.extra_attrs: config.extra_attrs["attn_layers"] = {} + suffix = 0 + # Makes sure there is no duplicate attention layer identifier. + while self.layer_idx_str in config.extra_attrs["attn_layers"]: + self.layer_idx_str = str(layer_idx) + f"_{suffix}" + suffix += 1 config.extra_attrs["attn_layers"][self.layer_idx_str] = weakref.ref( self) self.register_to_config = True @@ -358,7 +386,6 @@ def __init__( ) self.support_fused_qkv = self.attn.support_fused_qkv() - self.support_nvfp4_output = self.attn.support_nvfp4_output() if not config.skip_create_weights_in_init: self.create_weights() @@ -387,20 +414,22 @@ def convert_qkv(self, q, k, v): q, k, v = qkv, None, None return q, k, v - def create_output(self, q: torch.Tensor): - num_tokens = q.shape[0] - hidden_size = self.o_proj.in_features - out_dtype = q.dtype - - if self.attn_backend == "TRTLLM": - # Don't use FP8 output if o_proj has pre_quant_scale - keep BF16 for better precision - has_pre_quant_scale = getattr(self.o_proj, 'pre_quant_scale', - None) is not None - if self.has_quant_scale and not has_pre_quant_scale and ( - self.attn.has_fp8_kv_cache or self.attn.has_fp4_kv_cache): - out_dtype = torch.float8_e4m3fn - output = q.new_empty([num_tokens, hidden_size], dtype=out_dtype) - return output + def _use_quantize_output(self): + has_awq_pre_quant_scale = hasattr( + self.o_proj, + 'pre_quant_scale') and self.o_proj.pre_quant_scale is not None + + return self.has_quant_scale and not self.attn_output_gate and not has_awq_pre_quant_scale + + def create_output(self, q: torch.Tensor, attn_metadata: AttentionMetadata, + mask_type: str): + # Attention is treated as mixed request by default. + return self.attn.create_output( + q, + is_quantize_output=self._use_quantize_output(), + metadata=attn_metadata, + attention_mask=mask_type, + is_gen_only=False) def _attn_impl( self, @@ -413,7 +442,6 @@ def _attn_impl( mrope_position_deltas: Optional[torch.Tensor], attention_window_size: Optional[int], attention_mask_data: Optional[torch.Tensor], - enable_attn_nvfp4_output: bool = True, output: Optional[torch.Tensor] = None, output_sf: Optional[torch.Tensor] = None, attention_sinks: Optional[torch.Tensor] = None, @@ -428,19 +456,10 @@ def _attn_impl( out_scale = None out_scale_sf = None - has_awq_pre_quant_scale = hasattr( - self.o_proj, - 'pre_quant_scale') and self.o_proj.pre_quant_scale is not None # Don't set out_scale if o_proj has pre_quant_scale - this prevents FP8/FP4 output # and keeps attention output in BF16 for better precision when applying pre_quant_scale - if self.has_quant_scale and not self.attn_output_gate and not has_awq_pre_quant_scale: + if self._use_quantize_output(): out_scale = self.o_proj.inv_input_scale - if has_awq_pre_quant_scale and enable_attn_nvfp4_output: - logger.warning_once( - "Disable attn nvfp4 output because o_proj has pre_quant_scale for AWQ.", - key="disable_attn_nvfp4_output_for_awq") - enable_attn_nvfp4_output = False - if self.o_proj.has_nvfp4 and self.support_nvfp4_output and enable_attn_nvfp4_output and not self.attn_output_gate: out_scale_sf = self.o_proj.input_scale kv_scales_sf = None @@ -471,7 +490,6 @@ def _attn_impl( mrope_config=mrope_config, attention_window_size=attention_window_size, attention_mask_data=attention_mask_data, - enable_attn_nvfp4_output=enable_attn_nvfp4_output, output=output[:num_tokens, :] if output is not None else None, output_sf=output_sf, attention_sinks=attention_sinks) @@ -510,7 +528,10 @@ def forward_impl( and is_torch_compiling()) if use_custom_inplace_op: - output = self.create_output(q) + outputs = create_attn_outputs(q, attention_mask, self.layer_idx_str) + assert len(outputs) == 1 or len(outputs) == 2 + output = outputs[0] + output_sf = outputs[1] if len(outputs) == 2 else None attn_custom_op_inplace( q, k, @@ -523,6 +544,7 @@ def forward_impl( attention_sinks, self.layer_idx_str, output, + output_sf, ) else: output, output_sf = self._attn_impl(q, @@ -535,8 +557,8 @@ def forward_impl( attention_window_size, attention_mask_data, attention_sinks=attention_sinks) - if output_sf is not None: - output = Fp4QuantizedTensor(output, output_sf) + if output_sf is not None: + output = Fp4QuantizedTensor(output, output_sf) return output diff --git a/tensorrt_llm/_utils.py b/tensorrt_llm/_utils.py index 549575ce8f17..86ebaef371d8 100644 --- a/tensorrt_llm/_utils.py +++ b/tensorrt_llm/_utils.py @@ -1,4 +1,4 @@ -# SPDX-FileCopyrightText: Copyright (c) 2022-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-FileCopyrightText: Copyright (c) 2022-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 # # Licensed under the Apache License, Version 2.0 (the "License"); @@ -430,6 +430,7 @@ def torch_dtype_to_binding(dtype): torch.qint8: "|u1", torch.bool: "|b1", torch.bfloat16: "