diff --git a/tpu_inference/envs.py b/tpu_inference/envs.py index 5cf956f7de..4adeb1724f 100644 --- a/tpu_inference/envs.py +++ b/tpu_inference/envs.py @@ -82,6 +82,7 @@ MOE_ROUTE_PADDING_TO_EXPERT0: bool = False VLLM_TPU_BUCKET_PADDING_GAP: int = 0 VLLM_INCREMENTAL_FP8_LOADING: bool = False + VLLM_INCREMENTAL_MXFP4_LOADING: bool = False TPU_MESH_SORT_BY_COORDS: bool = False @@ -494,6 +495,10 @@ def _get_int_list_env() -> list[int]: # when initializing large FP8 models on smaller RAM TPUs such as TPU8i. "VLLM_INCREMENTAL_FP8_LOADING": env_bool("VLLM_INCREMENTAL_FP8_LOADING", default=False), + # Controls whether MXFP4 MoE layers perform incremental weight + # loading, sharding, and immediate host RAM cleanup. + "VLLM_INCREMENTAL_MXFP4_LOADING": + env_bool("VLLM_INCREMENTAL_MXFP4_LOADING", default=False), } diff --git a/tpu_inference/layers/vllm/quantization/base.py b/tpu_inference/layers/vllm/quantization/base.py index 89d6f5d115..13f951f4cd 100644 --- a/tpu_inference/layers/vllm/quantization/base.py +++ b/tpu_inference/layers/vllm/quantization/base.py @@ -13,14 +13,67 @@ # limitations under the License. from abc import ABC, abstractmethod +import ctypes +import ctypes.util +import gc +from typing import Optional +import jax import torch from vllm.logger import init_logger from vllm.model_executor.layers import linear as vllm_linear +from tpu_inference import envs + logger = init_logger(__name__) +def _free_torch_storage(tensor: Optional[torch.Tensor]) -> None: + """Safely frees the underlying CPU memory storage of a PyTorch tensor. + + Tries `untyped_storage().resize_(0)` first, with fallback to `set_(torch.storage.UntypedStorage())` + for 0-dim scalars or float8 dtypes that cannot be resized in-place. + """ + if tensor is None: + return + try: + tensor.untyped_storage().resize_(0) + except Exception: + try: + tensor.set_(torch.storage.UntypedStorage()) + except Exception: + pass + + +def _release_host_memory() -> None: + """Frees CPU host memory and trims malloc arena if incremental loading is enabled.""" + if not (getattr(envs, "VLLM_INCREMENTAL_FP8_LOADING", False) or getattr(envs, "VLLM_INCREMENTAL_MXFP4_LOADING", False)): + return + gc.collect() + jax.effects_barrier() + try: + libc_name = ctypes.util.find_library("c") + if libc_name: + ctypes.CDLL(libc_name).malloc_trim(0) + except Exception as e: + logger.debug(f"malloc_trim failed: {e}") + + +def _log_memory_stats(layer_name: str = "") -> None: + try: + import psutil, resource + proc = psutil.Process() + rss_gb = proc.memory_info().rss / (1024 ** 3) + max_rss_gb = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / (1024 ** 2) + print( + f"[RAM Trace] Layer {layer_name} sharded & freed | " + f"Process RSS: {rss_gb:.2f} GB | Peak RSS: {max_rss_gb:.2f} GB", + flush=True, + ) + except Exception as e: + print(f"[RAM Trace Error] {e}", flush=True) + + class VllmQuantizationMethod(ABC): def maybe_process_linear_weights( diff --git a/tpu_inference/layers/vllm/quantization/fp8.py b/tpu_inference/layers/vllm/quantization/fp8.py index f83754f0a9..819232fc0f 100644 --- a/tpu_inference/layers/vllm/quantization/fp8.py +++ b/tpu_inference/layers/vllm/quantization/fp8.py @@ -51,7 +51,8 @@ select_moe_backend_from_fused_moe_config, vllm_moe_apply) from tpu_inference.layers.vllm.process_weights.cleanup_sharding import \ _tensor_is_in_cpu -from tpu_inference.layers.vllm.quantization.base import VllmQuantizationMethod +from tpu_inference.layers.vllm.quantization.base import ( + VllmQuantizationMethod, _log_memory_stats) from tpu_inference.layers.vllm.quantization.configs import ( VllmQuantConfig, VllmQuantLinearConfig) from tpu_inference.layers.vllm.quantization.unquantized import ( @@ -324,6 +325,7 @@ def process_weights_after_loading(self, layer: torch.nn.Module) -> None: layer.bias = to_parameter_list(weights.bias) _release_host_memory() + _log_memory_stats(layer_name=getattr(layer, "_module_name", getattr(layer, "prefix", str(type(layer))))) def apply(self, layer: torch.nn.Module, @@ -473,11 +475,14 @@ def process_weights_after_loading(self, layer: torch.nn.Module) -> None: del w13_weight, w2_weight, w13_weight_scale, w2_weight_scale, input_weights - weights = torch_view( - shard_moe_weights(weights, self.moe_backend, self.mesh)) + sharded = shard_moe_weights(weights, self.moe_backend, self.mesh) + del weights + + tv_weights = torch_view(sharded) + del sharded - layer.w13_weight = Parameter(weights.w13_weight, requires_grad=False) - layer.w2_weight = Parameter(weights.w2_weight, requires_grad=False) + layer.w13_weight = Parameter(tv_weights.w13_weight, requires_grad=False) + layer.w2_weight = Parameter(tv_weights.w2_weight, requires_grad=False) # Use setattr to dynamically assign the correct scale parameter name # based on the quantization type. vLLM uses 'weight_scale_inv' for @@ -485,14 +490,17 @@ def process_weights_after_loading(self, layer: torch.nn.Module) -> None: setattr( layer, scale_w13_name, - Parameter(weights.w13_weight_scale, requires_grad=False), + Parameter(tv_weights.w13_weight_scale, requires_grad=False), ) setattr( layer, scale_w2_name, - Parameter(weights.w2_weight_scale, requires_grad=False), + Parameter(tv_weights.w2_weight_scale, requires_grad=False), ) + del tv_weights + _release_host_memory() + _log_memory_stats(layer_name=getattr(layer, "_module_name", getattr(layer, "prefix", str(type(layer))))) def apply_monolithic( self, diff --git a/tpu_inference/layers/vllm/quantization/mxfp4.py b/tpu_inference/layers/vllm/quantization/mxfp4.py index 308e1d2bd9..6e941c2f12 100644 --- a/tpu_inference/layers/vllm/quantization/mxfp4.py +++ b/tpu_inference/layers/vllm/quantization/mxfp4.py @@ -12,12 +12,13 @@ # See the License for the specific language governing permissions and # limitations under the License. -from typing import Optional +import functools +from typing import Any, Optional import jax import jax.numpy as jnp import torch -from jax.sharding import Mesh, PartitionSpec +from jax.sharding import Mesh, NamedSharding, PartitionSpec from torch.nn.parameter import Parameter from torchax.interop import jax_view, torch_view from vllm.model_executor.layers.attention import Attention @@ -38,6 +39,7 @@ from vllm.model_executor.layers.quantization.utils.quant_utils import \ is_layer_skipped +from tpu_inference import envs from tpu_inference.layers.common.moe import \ FusedMoEMethodBase as TpuFusedMoEMethodBase from tpu_inference.layers.common.process_weights.moe_weights import ( @@ -49,11 +51,16 @@ from tpu_inference.layers.common.sharding import ShardingAxisName from tpu_inference.layers.vllm.interface.moe import ( select_moe_backend_from_fused_moe_config, vllm_moe_apply) +from tpu_inference.layers.vllm.process_weights.cleanup_sharding import \ + _tensor_is_in_cpu +from tpu_inference.layers.vllm.quantization.base import ( + VllmQuantizationMethod, _free_torch_storage, _log_memory_stats, + _release_host_memory) from tpu_inference.layers.vllm.quantization.configs import VllmQuantConfig -from tpu_inference.layers.vllm.quantization.unquantized import \ - VllmUnquantizedLinearMethod +from tpu_inference.layers.vllm.quantization.unquantized import ( + VllmUnquantizedLinearMethod, _load_weight_for_layer) from tpu_inference.logger import init_logger -from tpu_inference.utils import get_mesh_shape_product, t2j +from tpu_inference.utils import get_mesh_shape_product, t2j, to_jax_dtype P = PartitionSpec @@ -91,7 +98,57 @@ def get_quant_method(self, layer: torch.nn.Module, return None -class VllmMxfp4MoEMethod(Mxfp4MoEMethod, FusedMoEMethodBase): +@functools.partial( + jax.jit, + static_argnames=( + "moe_backend", + "w13_reorder_size", + "w13_interleave", + "desired_quant_dtype", + "requant_block_size", + ), +) +def _process_mxfp4_moe_weights( + w13_weight: jax.Array, + w13_weight_scale: jax.Array, + w13_bias: jax.Array | None, + w2_weight: jax.Array, + w2_weight_scale: jax.Array, + w2_bias: jax.Array | None, + moe_backend: Any, + w13_reorder_size: int, + w13_interleave: bool, + desired_quant_dtype: Any, + requant_block_size: int, +) -> FusedMoEWeights: + # Dequantize fp4 weights into fp32. + w13_weight = dequantize_tensor_from_mxfp4_packed( + w13_weight, w13_weight_scale, 2, jnp.float32) + w2_weight = dequantize_tensor_from_mxfp4_packed( + w2_weight, w2_weight_scale, 2, jnp.float32) + + weights = quantize_moe_weights( + FusedMoEWeights( + w13_weight=w13_weight, + w13_weight_scale=None, + w13_bias=w13_bias, + w2_weight=w2_weight, + w2_weight_scale=None, + w2_bias=w2_bias, + ), + desired_quant_dtype, + requant_block_size, + w13_interleave=w13_interleave, + ) + return process_moe_weights( + weights, + moe_backend=moe_backend, + w13_reorder_size=w13_reorder_size, + w13_interleave=w13_interleave, + ) + + +class VllmMxfp4MoEMethod(Mxfp4MoEMethod, FusedMoEMethodBase, VllmQuantizationMethod): def __init__( self, @@ -129,78 +186,121 @@ def get_fused_moe_quant_config( def is_monolithic(self) -> bool: return True + def maybe_process_weights(self, layer: torch.nn.Module, param_name: str, + args, kwargs): + """Check if all weights are loaded for the MoE layer. + + If so, process and shard the weights. + """ + if not (getattr(envs, "VLLM_INCREMENTAL_FP8_LOADING", False) or getattr(envs, "VLLM_INCREMENTAL_MXFP4_LOADING", False)): + return + + logger.info_once( + "[mxfp4-incremental] VLLM_INCREMENTAL_MXFP4_LOADING is enabled") + + expert_id = kwargs.get("expert_id") + shard_id = kwargs.get("shard_id") + assert expert_id is not None, "Expecting expert_id argument" + assert shard_id is not None, "Expecting shard_id argument" + + layer._loaded_weights.add((param_name, expert_id, shard_id)) + + num_experts = getattr(layer, "global_num_experts", getattr(layer, "num_experts", None)) + assert num_experts is not None, "Layer must have global_num_experts or num_experts" + expected_shards = 6 * num_experts + + if len(layer._loaded_weights) >= expected_shards: + logger.debug( + f"[mxfp4-incremental] Start sharding weights for MoE layer {type(layer)}" + ) + self.process_weights_after_loading(layer) + logger.debug( + f"[mxfp4-incremental] Complete sharding weights for MoE layer {type(layer)}" + ) + def process_weights_after_loading(self, layer: torch.nn.Module): + if not hasattr(layer, "w13_weight") or not _tensor_is_in_cpu(layer.w13_weight): + return assert isinstance(layer, RoutedExperts) has_bias = layer.moe_config.has_bias - w13_weight = t2j(layer.w13_weight, use_dlpack=False) - w13_weight_scale = t2j(layer.w13_weight_scale, use_dlpack=False) - w13_bias = t2j(layer.w13_bias, use_dlpack=False) if has_bias else None - - w2_weight = t2j(layer.w2_weight, use_dlpack=False) - w2_weight_scale = t2j(layer.w2_weight_scale, use_dlpack=False) - w2_bias = t2j(layer.w2_bias, use_dlpack=False) if has_bias else None - - @jax.jit - def process_mxfp4_moe_weights( - w13_weight: jax.Array, - w13_weight_scale: jax.Array, - w13_bias: jax.Array | None, - w2_weight: jax.Array, - w2_weight_scale: jax.Array, - w2_bias: jax.Array | None, - ) -> FusedMoEWeights: - # Dequantize fp4 weights into fp32. - w13_weight = dequantize_tensor_from_mxfp4_packed( - w13_weight, w13_weight_scale, 2, jnp.float32) - w2_weight = dequantize_tensor_from_mxfp4_packed( - w2_weight, w2_weight_scale, 2, jnp.float32) - w13_interleave = layer.activation == MoEActivation.SWIGLUOAI - w13_reorder_size = get_mesh_shape_product( - self.mesh, ShardingAxisName.MLP_TENSOR) - - weights = quantize_moe_weights( - FusedMoEWeights( - w13_weight=w13_weight, - w13_weight_scale=None, - w13_bias=w13_bias, - w2_weight=w2_weight, - w2_weight_scale=None, - w2_bias=w2_bias, - ), - jnp.float4_e2m1fn, - MXFP4_REQUANTIZED_BLOCK_SIZE, - w13_interleave=w13_interleave, - ) - return process_moe_weights( - weights, - moe_backend=self.moe_backend, - w13_reorder_size=w13_reorder_size, - w13_interleave=w13_interleave, - ) + ep_sharding = NamedSharding(self.mesh, P(ShardingAxisName.EXPERT)) + + p_w13_weight = layer.w13_weight + p_w13_scale = layer.w13_weight_scale + p_w2_weight = layer.w2_weight + p_w2_scale = layer.w2_weight_scale + + w13_weight = _load_weight_for_layer(layer, "w13_weight", ep_sharding) + w13_weight_scale = _load_weight_for_layer(layer, "w13_weight_scale", ep_sharding) + w2_weight = _load_weight_for_layer(layer, "w2_weight", ep_sharding) + w2_weight_scale = _load_weight_for_layer(layer, "w2_weight_scale", ep_sharding) + + p_w13_bias = getattr(layer, "w13_bias", None) if has_bias else None + p_w2_bias = getattr(layer, "w2_bias", None) if has_bias else None + w13_bias = _load_weight_for_layer(layer, "w13_bias", ep_sharding) if has_bias else None + w2_bias = _load_weight_for_layer(layer, "w2_bias", ep_sharding) if has_bias else None - weights = process_mxfp4_moe_weights( + w13_interleave = layer.activation == MoEActivation.SWIGLUOAI + w13_reorder_size = get_mesh_shape_product( + self.mesh, ShardingAxisName.MLP_TENSOR) + + desired_quant_dtype = to_jax_dtype(envs.MOE_REQUANTIZE_WEIGHT_DTYPE) if envs.MOE_REQUANTIZE_WEIGHT_DTYPE else jnp.float4_e2m1fn + requant_block_size = int(envs.MOE_REQUANTIZE_BLOCK_SIZE) if envs.MOE_REQUANTIZE_BLOCK_SIZE else MXFP4_REQUANTIZED_BLOCK_SIZE + + weights = _process_mxfp4_moe_weights( w13_weight, w13_weight_scale, w13_bias, w2_weight, w2_weight_scale, w2_bias, + moe_backend=self.moe_backend, + w13_reorder_size=w13_reorder_size, + w13_interleave=w13_interleave, + desired_quant_dtype=desired_quant_dtype, + requant_block_size=requant_block_size, ) - weights = torch_view( - shard_moe_weights(weights, self.moe_backend, self.mesh)) - layer.w13_weight = Parameter(weights.w13_weight, requires_grad=False) - layer.w2_weight = Parameter(weights.w2_weight, requires_grad=False) + # Free CPU memory now that weights have been safely transferred to TPU + _free_torch_storage(p_w13_weight) + _free_torch_storage(p_w13_scale) + _free_torch_storage(p_w2_weight) + _free_torch_storage(p_w2_scale) + delattr(layer, "w13_weight") + delattr(layer, "w13_weight_scale") + delattr(layer, "w2_weight") + delattr(layer, "w2_weight_scale") + if has_bias: + _free_torch_storage(p_w13_bias) + _free_torch_storage(p_w2_bias) + delattr(layer, "w13_bias") + delattr(layer, "w2_bias") - layer.w13_weight_scale = Parameter(weights.w13_weight_scale, + del w13_weight, w13_weight_scale, w2_weight, w2_weight_scale, w13_bias, w2_bias + + sharded_weights = shard_moe_weights(weights, self.moe_backend, self.mesh) + del weights + + tv_weights = torch_view(sharded_weights) + del sharded_weights + + layer.w13_weight = Parameter(tv_weights.w13_weight, requires_grad=False) + layer.w2_weight = Parameter(tv_weights.w2_weight, requires_grad=False) + + layer.w13_weight_scale = Parameter(tv_weights.w13_weight_scale, requires_grad=False) - layer.w2_weight_scale = Parameter(weights.w2_weight_scale, + layer.w2_weight_scale = Parameter(tv_weights.w2_weight_scale, requires_grad=False) if has_bias: - layer.w13_bias = Parameter(weights.w13_bias, requires_grad=False) - layer.w2_bias = Parameter(weights.w2_bias, requires_grad=False) + layer.w13_bias = Parameter(tv_weights.w13_bias, requires_grad=False) + layer.w2_bias = Parameter(tv_weights.w2_bias, requires_grad=False) + + del tv_weights + + _release_host_memory() + _log_memory_stats(layer_name=getattr(layer, "_module_name", getattr(layer, "prefix", str(type(layer))))) def apply_monolithic( self, diff --git a/tpu_inference/layers/vllm/quantization/unquantized.py b/tpu_inference/layers/vllm/quantization/unquantized.py index 5aa6e6695b..a99da35119 100644 --- a/tpu_inference/layers/vllm/quantization/unquantized.py +++ b/tpu_inference/layers/vllm/quantization/unquantized.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import functools from typing import Any, Callable, Optional import jax @@ -51,7 +52,8 @@ select_moe_backend_from_fused_moe_config, vllm_moe_apply) from tpu_inference.layers.vllm.process_weights.cleanup_sharding import \ _tensor_is_in_cpu -from tpu_inference.layers.vllm.quantization.base import VllmQuantizationMethod +from tpu_inference.layers.vllm.quantization.base import ( + VllmQuantizationMethod, _log_memory_stats, _release_host_memory) from tpu_inference.layers.vllm.quantization.configs import ( VllmQuantConfig, VllmQuantLinearConfig) from tpu_inference.logger import init_logger @@ -197,6 +199,34 @@ def process_weights_after_loading(self, layer: torch.nn.Module) -> None: layer.bias = Parameter(torch_view(bias), requires_grad=False) +@functools.partial( + jax.jit, + static_argnames=( + "fused", + "output_sizes", + "reorder_size", + ), +) +def _process_unquantized_linear_weights( + weight: jax.Array, + bias: jax.Array | None, + fused: bool, + output_sizes: tuple[int, ...], + reorder_size: int, +) -> LinearWeights: + return process_linear_weights( + LinearWeights( + weight=weight, + weight_scale=None, + zero_point=None, + bias=bias, + ), + fused=fused, + output_sizes=list(output_sizes), + reorder_size=reorder_size, + ) + + class VllmUnquantizedLinearMethod(vllm_linear.UnquantizedLinearMethod, common_unquantized.UnquantizedLinearMethod, VllmQuantizationMethod): @@ -241,39 +271,38 @@ def process_weights_after_loading(self, layer: torch.nn.Module) -> None: else: bias = None - @jax.jit - def process_unquantized_linear_weights( - weight: jax.Array, - bias: jax.Array | None, - ) -> LinearWeights: - return process_linear_weights( - LinearWeights( - weight=weight, - weight_scale=None, - zero_point=None, - bias=bias, - ), - fused=self.linear_config.fuse_matmuls, - output_sizes=self.linear_config.output_sizes, - reorder_size=self.linear_config.n_shards, - ) + weights = _process_unquantized_linear_weights( + weight, + bias, + fused=self.linear_config.fuse_matmuls, + output_sizes=tuple(self.linear_config.output_sizes), + reorder_size=self.linear_config.n_shards, + ) + del weight, bias + + sharded = shard_linear_weights( + weights, + mesh=self.linear_config.mesh, + weight_p_spec=self.linear_config.weight_sharding, + bias_p_spec=self.linear_config.bias_sharding, + ) + del weights + + tv_weights = torch_view(sharded) + del sharded - weights = process_unquantized_linear_weights(weight, bias) - weights = torch_view( - shard_linear_weights( - weights, - mesh=self.linear_config.mesh, - weight_p_spec=self.linear_config.weight_sharding, - bias_p_spec=self.linear_config.bias_sharding, - )) if self.linear_config.fuse_matmuls: - layer.weight = Parameter(weights.weight, requires_grad=False) - if bias is not None: - layer.bias = Parameter(weights.bias, requires_grad=False) + layer.weight = Parameter(tv_weights.weight, requires_grad=False) + if tv_weights.bias is not None: + layer.bias = Parameter(tv_weights.bias, requires_grad=False) else: - layer.weight = to_parameter_list(weights.weight) - if bias is not None: - layer.bias = to_parameter_list(weights.bias) + layer.weight = to_parameter_list(tv_weights.weight) + if tv_weights.bias is not None: + layer.bias = to_parameter_list(tv_weights.bias) + del tv_weights + + _release_host_memory() + _log_memory_stats(layer_name=getattr(layer, "_module_name", getattr(layer, "prefix", str(type(layer))))) def apply(self, layer: torch.nn.Module, @@ -384,19 +413,25 @@ def process_weights_after_loading(self, layer: torch.nn.Module) -> None: del w13_weight, w2_weight, w13_bias, w2_bias - weights = torch_view( - shard_moe_weights(weights, self.moe_backend, self.mesh)) - layer.w13_weight = Parameter(weights.w13_weight, requires_grad=False) - layer.w2_weight = Parameter(weights.w2_weight, requires_grad=False) + sharded = shard_moe_weights(weights, self.moe_backend, self.mesh) + del weights + + tv_weights = torch_view(sharded) + del sharded + + layer.w13_weight = Parameter(tv_weights.w13_weight, requires_grad=False) + layer.w2_weight = Parameter(tv_weights.w2_weight, requires_grad=False) if self.moe.has_bias: - layer.w13_bias = Parameter(weights.w13_bias, requires_grad=False) - layer.w2_bias = Parameter(weights.w2_bias, requires_grad=False) + layer.w13_bias = Parameter(tv_weights.w13_bias, requires_grad=False) + layer.w2_bias = Parameter(tv_weights.w2_bias, requires_grad=False) + del tv_weights # Force JAX to release intermediate buffers before processing the next # layer. Without this barrier, async dispatch can keep old weight # buffers alive across layers, accumulating until OOM. - jax.effects_barrier() + _release_host_memory() + _log_memory_stats(layer_name=getattr(layer, "_module_name", getattr(layer, "prefix", str(type(layer))))) def apply_monolithic( self, diff --git a/tpu_inference/models/jax/jax_intermediate_tensor.py b/tpu_inference/models/jax/jax_intermediate_tensor.py index 3ac64cc5dd..f92cc7a19c 100644 --- a/tpu_inference/models/jax/jax_intermediate_tensor.py +++ b/tpu_inference/models/jax/jax_intermediate_tensor.py @@ -59,6 +59,8 @@ def tree_unflatten(cls, aux_data, children): @classmethod def from_torch(cls, torch_obj: IntermediateTensors): + if not hasattr(torch_obj, 'tensors'): + return cls({'hidden_states': jax_view(torch_obj)}) kv_connector_output = getattr(torch_obj, 'kv_connector_output', None) jax_tensors = {k: jax_view(v) for k, v in torch_obj.tensors.items()} return cls(jax_tensors, kv_connector_output) diff --git a/tpu_inference/models/vllm/experimental/__init__.py b/tpu_inference/models/vllm/experimental/__init__.py index a22732f439..600ecc8271 100644 --- a/tpu_inference/models/vllm/experimental/__init__.py +++ b/tpu_inference/models/vllm/experimental/__init__.py @@ -28,7 +28,10 @@ def register_models(): Called from the ``vllm.general_plugins`` entrypoint so it runs at vLLM startup, before any model is resolved from its architecture name. """ - from vllm import ModelRegistry + try: + from vllm import ModelRegistry + except ImportError: + from vllm.model_executor.models import ModelRegistry for arch, model_cls in _TPU_VLLM_MODELS.items(): ModelRegistry.register_model(arch, model_cls) diff --git a/tpu_inference/models/vllm/experimental/deepseek_v4.py b/tpu_inference/models/vllm/experimental/deepseek_v4.py index eb321fed68..5d9b234ddd 100644 --- a/tpu_inference/models/vllm/experimental/deepseek_v4.py +++ b/tpu_inference/models/vllm/experimental/deepseek_v4.py @@ -8,7 +8,8 @@ import torch import torch.nn as nn from vllm.config import VllmConfig -from vllm.distributed import (get_pp_group, get_tensor_model_parallel_rank, +from tpu_inference.distributed.jax_parallel_state import get_pp_group +from vllm.distributed import (get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size) from vllm.model_executor.layers.activation import (SiluAndMul, SiluAndMulWithClamp) diff --git a/tpu_inference/models/vllm/vllm_model_loader.py b/tpu_inference/models/vllm/vllm_model_loader.py index 8da8ffc685..62b14e2fab 100644 --- a/tpu_inference/models/vllm/vllm_model_loader.py +++ b/tpu_inference/models/vllm/vllm_model_loader.py @@ -26,7 +26,9 @@ initialize_model, process_weights_after_loading) from vllm.utils.torch_utils import set_default_torch_dtype -from tpu_inference.layers.vllm.quantization.base import VllmQuantizationMethod +from tpu_inference import envs +from tpu_inference.layers.vllm.quantization.base import ( + VllmQuantizationMethod, _free_torch_storage, _release_host_memory) def attach_incremental_weight_loader(model: torch.nn.Module) -> None: @@ -34,12 +36,24 @@ def attach_incremental_weight_loader(model: torch.nn.Module) -> None: Traverses the model and overrides the weight_loader of each parameter to support incremental loading. This allows processing and sharding of weights after all weights for a module have been loaded. """ + is_incremental_enabled = ( + getattr(envs, "VLLM_INCREMENTAL_FP8_LOADING", False) or + getattr(envs, "VLLM_INCREMENTAL_MXFP4_LOADING", False) + ) def create_weight_loader(layer, original_loader, layer_name, param_name): def weight_loader_wrapper(param: torch.nn.Parameter, loaded_weight: torch.Tensor, *args, **kwargs): + # If parameter storage was lazily freed upfront, reallocate zero-initialized buffer on first shard. + if hasattr(param, "_full_shape") and param.untyped_storage().size() == 0: + param.data = torch.zeros( + param._full_shape, + dtype=param._full_dtype, + device=param._full_device, + ) + # Loading the weight res = original_loader(param, loaded_weight, *args, **kwargs) @@ -59,6 +73,10 @@ def weight_loader_wrapper(param: torch.nn.Parameter, # Weight loader will be invoked multiple times for module. In order to determine when all the weights are loaded, # we need to keep track of the loaded weights for each module. module._loaded_weights = set() + module._module_name = name + quant_method = getattr(module, "quant_method", None) + is_incremental_layer = isinstance(quant_method, VllmQuantizationMethod) + for param_name, param in module.named_parameters(recurse=False): # Omit parameters that do not have a weight_loader original_loader = getattr(param, "weight_loader", None) @@ -69,6 +87,15 @@ def weight_loader_wrapper(param: torch.nn.Parameter, create_weight_loader(module, original_loader, name, param_name)) + if is_incremental_enabled and is_incremental_layer: + param._full_shape = tuple(param.shape) + param._full_dtype = param.dtype + param._full_device = param.device + _free_torch_storage(param) + + if is_incremental_enabled: + _release_host_memory() + @register_model_loader("tpu_streaming_loader") class IncrementalModelLoader(DefaultModelLoader): diff --git a/tpu_inference/runner/compilation_manager.py b/tpu_inference/runner/compilation_manager.py index 3a35d197f7..6cd3dfc061 100644 --- a/tpu_inference/runner/compilation_manager.py +++ b/tpu_inference/runner/compilation_manager.py @@ -304,6 +304,17 @@ def _finalize_compilation(self) -> None: except (RuntimeError, ValueError): pass self._prev_stack_size = None + try: + from flax import nnx + from tpu_inference.utils import device_array + rng_key = nnx.Rngs(jax.random.key(self.runner.model_config.seed)).params() + self.runner.rng_params_for_sampling = device_array( + self.runner.mesh, + rng_key, + sharding=NamedSharding(self.runner.mesh, PartitionSpec())) + logger.info("Successfully re-initialized rng_params_for_sampling after compilation.") + except Exception as e: + logger.warning(f"Failed to reset rng_params_for_sampling: {e}") def _precompile_input_embeddings_merger(self) -> None: for num_tokens in self.runner.num_tokens_paddings: @@ -572,7 +583,7 @@ def _compile_one(input_padding: int, input_sharding: NamedSharding, padded_token_in_tpu_pre_next_tokens_indices, next_tokens, placeholder_num, - compile_only=False, + compile_only=True, num_tokens=input_padding, next_tokens_size=next_tokens_size, ) @@ -700,19 +711,37 @@ def _precompile_backbone_text_only(self) -> None: sharding = NamedSharding( self.runner.mesh, PartitionSpec(ShardingAxisName.ATTN_DATA, None)) - hidden_states = self._create_dummy_tensor( - (num_tokens, hidden_size), - jnp.bfloat16, - sharding=sharding) - residual = self._create_dummy_tensor( - (num_tokens, hidden_size), - jnp.bfloat16, - sharding=sharding) - intermediate_tensors = JaxIntermediateTensors( - tensors={ - "hidden_states": hidden_states, - "residual": residual - }) + hf_conf = self.runner.vllm_config.model_config.hf_config + hc_mult = getattr(hf_conf, "hc_mult", None) + if hc_mult: + hs_shape = (num_tokens, hc_mult, hidden_size) + hs_sharding = NamedSharding( + self.runner.mesh, + PartitionSpec(ShardingAxisName.ATTN_DATA, None, None)) + hidden_states = self._create_dummy_tensor( + hs_shape, + jnp.bfloat16, + sharding=hs_sharding) + intermediate_tensors = JaxIntermediateTensors( + tensors={ + "hidden_states": hidden_states, + }) + else: + hs_shape = (num_tokens, hidden_size) + hs_sharding = sharding + hidden_states = self._create_dummy_tensor( + hs_shape, + jnp.bfloat16, + sharding=hs_sharding) + residual = self._create_dummy_tensor( + (num_tokens, hidden_size), + jnp.bfloat16, + sharding=sharding) + intermediate_tensors = JaxIntermediateTensors( + tensors={ + "hidden_states": hidden_states, + "residual": residual + }) for _cache_pages in self._pcp_cache_page_buckets(): self._precompile_backbone_helper( f"worker{self.runner.rank} backbone", @@ -773,19 +802,37 @@ def _precompile_backbone_with_inputs_embeds(self) -> None: is_first_rank = self.runner.is_first_rank is_last_rank = self.runner.is_last_rank if not is_first_rank: - hidden_states = self._create_dummy_tensor( - (num_tokens, hidden_size), - jnp.bfloat16, - sharding=sharding) - residual = self._create_dummy_tensor( - (num_tokens, hidden_size), - jnp.bfloat16, - sharding=sharding) - intermediate_tensors = JaxIntermediateTensors( - tensors={ - "hidden_states": hidden_states, - "residual": residual - }) + hf_conf = self.runner.vllm_config.model_config.hf_config + hc_mult = getattr(hf_conf, "hc_mult", None) + if hc_mult: + hs_shape = (num_tokens, hc_mult, hidden_size) + hs_sharding = NamedSharding( + self.runner.mesh, + PartitionSpec(ShardingAxisName.ATTN_DATA, None, None)) + hidden_states = self._create_dummy_tensor( + hs_shape, + jnp.bfloat16, + sharding=hs_sharding) + intermediate_tensors = JaxIntermediateTensors( + tensors={ + "hidden_states": hidden_states, + }) + else: + hs_shape = (num_tokens, hidden_size) + hs_sharding = sharding + hidden_states = self._create_dummy_tensor( + hs_shape, + jnp.bfloat16, + sharding=hs_sharding) + residual = self._create_dummy_tensor( + (num_tokens, hidden_size), + jnp.bfloat16, + sharding=sharding) + intermediate_tensors = JaxIntermediateTensors( + tensors={ + "hidden_states": hidden_states, + "residual": residual + }) else: intermediate_tensors = None self._precompile_backbone_helper( @@ -966,8 +1013,9 @@ def _precompile_sampling(self) -> None: # function. sampling_metadata_sharding = NamedSharding( self.runner.mesh, PartitionSpec(ShardingAxisName.ATTN_DATA)) + from tpu_inference.utils import to_jax_dtype logits = self._create_dummy_tensor((num_reqs, hsize), - jnp.float32, + to_jax_dtype(self.runner.dtype), sharding=logits_sharding) for do_sampling in (True, False): for logprobs in (True, False): @@ -1009,7 +1057,7 @@ def _precompile_sampling(self) -> None: self.runner.mesh, logits, sampling_metadata, - compile_only=False, + compile_only=True, num_reqs=num_reqs, do_sampling=do_sampling, logprobs=logprobs, diff --git a/tpu_inference/runner/persistent_batch_manager.py b/tpu_inference/runner/persistent_batch_manager.py index 05e827586e..293948da8a 100644 --- a/tpu_inference/runner/persistent_batch_manager.py +++ b/tpu_inference/runner/persistent_batch_manager.py @@ -223,20 +223,17 @@ def update_states(self, scheduler_output: "VllmSchedulerOutput", req_state.num_computed_tokens = num_computed_tokens req_index = self.input_batch.req_id_to_index.get(req_id) + new_token_ids = req_data.new_token_ids[i] if req_data.new_token_ids else [] if not self.is_last_rank: - # When using PP, the scheduler sends the sampled tokens back, - # because there's no direct communication between the first- - # stage worker and the last-stage worker. - new_token_ids = req_data.new_token_ids[i] - # Add the sampled token(s) from the previous step (if any). - # This doesn't include "unverified" tokens like spec tokens. - num_new_tokens = (num_computed_tokens + len(new_token_ids) - - req_state.num_tokens) - if num_new_tokens == 1: - req_state.output_token_ids.append(new_token_ids[-1]) - elif num_new_tokens > 0: - req_state.output_token_ids.extend( - new_token_ids[-num_new_tokens:]) + # Only append output tokens if new_token_ids is non-empty + if new_token_ids: + num_new_tokens = (num_computed_tokens + len(new_token_ids) - + req_state.num_tokens) + if num_new_tokens == 1: + req_state.output_token_ids.append(new_token_ids[-1]) + elif num_new_tokens > 0: + req_state.output_token_ids.extend( + new_token_ids[-num_new_tokens:]) elif num_output_tokens < len(req_state.output_token_ids): del req_state.output_token_ids[num_output_tokens:] if req_index is not None: @@ -272,9 +269,8 @@ def update_states(self, scheduler_output: "VllmSchedulerOutput", self.input_batch.block_table.append_row( new_block_ids, req_index) - # For the last rank, we don't need to update the token_ids_cpu - # because the sampled tokens are already cached. - if not self.is_last_rank: + # For non-last ranks, update token_ids_cpu and token counts with the new token IDs from scheduler. + if not self.is_last_rank and new_token_ids: start_token_index = num_computed_tokens end_token_index = num_computed_tokens + len(new_token_ids) self.input_batch.token_ids_cpu[ diff --git a/tpu_inference/runner/tpu_runner.py b/tpu_inference/runner/tpu_runner.py index b2e96075c2..7c87d7a6eb 100644 --- a/tpu_inference/runner/tpu_runner.py +++ b/tpu_inference/runner/tpu_runner.py @@ -376,7 +376,8 @@ def __init__(self, scheduler_output: Optional["VllmSchedulerOutput"] = None, req_ids_dp: Optional[Dict] = None, padded_num_scheduled_tokens_per_dp_rank: int = 0, - runner=None): + runner=None, + valid_sampled_token_ids: Optional[List] = None): self._model_runner_output = model_runner_output self._next_tokens = next_tokens self._num_reqs = num_reqs @@ -391,6 +392,7 @@ def __init__(self, self._req_ids_dp = req_ids_dp self._padded_num_scheduled_tokens_per_dp_rank = padded_num_scheduled_tokens_per_dp_rank self._runner = runner + self._valid_sampled_token_ids = valid_sampled_token_ids self._is_continue_decode = False self._actual_steps_future = None @@ -443,12 +445,19 @@ def get_output(self) -> ModelRunnerOutput: if getattr(self, "_is_continue_decode", False): return self._get_continue_decode_output() - valid_sampled_token_ids = runner_utils.host_extract_sampled_tokens( - self._runner, self._spec_decode_metadata, self._next_tokens, - self.logits_indices_selector, - self._discard_sampled_tokens_req_indices, self._num_reqs) + if getattr(self, "_valid_sampled_token_ids", None) is not None: + valid_sampled_token_ids = self._valid_sampled_token_ids + elif self._runner and getattr(self._runner, "_pre_async_results", None) is not None and getattr(self._runner._pre_async_results, "valid_sampled_token_ids", None) is not None: + valid_sampled_token_ids = self._runner._pre_async_results.valid_sampled_token_ids + else: + valid_sampled_token_ids = runner_utils.host_extract_sampled_tokens( + self._runner, self._spec_decode_metadata, self._next_tokens, + self.logits_indices_selector, + self._discard_sampled_tokens_req_indices, self._num_reqs) self._model_runner_output.sampled_token_ids = valid_sampled_token_ids + if self._runner and getattr(self._runner, "_pre_async_results", None) is not None: + self._runner._pre_async_results.valid_sampled_token_ids = valid_sampled_token_ids if self._logprobs_tensors is not None: # Use materialize to ensure logprobs are ready on host when we return async results @@ -500,6 +509,7 @@ class AsyncPreResults: spec_decode_num_rejected_tokens: Optional[ jax.Array] = None # [max_num_reqs] spec_decode_metadata: Optional[SpecDecodeMetadata] = None + valid_sampled_token_ids: Optional[list] = None # For continue decode async scheduling is_continue_decode: bool = False @@ -533,7 +543,7 @@ class ExecuteModelState: padded_num_scheduled_tokens_per_dp_rank: int = 0 -@jax.jit(donate_argnums=(0, 1, 2)) +@jax.jit def _substitute_placeholder_token( input_ids: jax.Array, token_in_tpu_cur_input_indices: jax.Array, token_in_tpu_pre_next_tokens_indices: jax.Array, @@ -1453,10 +1463,13 @@ def _modify_prev_results(self): req_id] = actual_len return - valid_sampled_token_ids = runner_utils.host_extract_sampled_tokens( - self, pre_spec_decode_metadata, pre_next_tokens, - pre_logits_indices_selector, - pre_discard_sampled_tokens_req_indices, pre_num_reqs) + if getattr(self._pre_async_results, "valid_sampled_token_ids", None) is not None: + valid_sampled_token_ids = self._pre_async_results.valid_sampled_token_ids + else: + valid_sampled_token_ids = runner_utils.host_extract_sampled_tokens( + self, pre_spec_decode_metadata, pre_next_tokens, + pre_logits_indices_selector, + pre_discard_sampled_tokens_req_indices, pre_num_reqs) # Append sampled tokens for pre_req_idx, req_state, _ in pre_request_seq_lens: @@ -1556,7 +1569,7 @@ def _execute_model( self.persistent_batch_manager.update_states( scheduler_output, self.get_mrope_input_positions_fn) if not scheduler_output.total_num_scheduled_tokens: - if self.scheduler_config.async_scheduling and self._pre_async_results is not None: + if self.scheduler_config.async_scheduling and self._pre_async_results is not None and self.parallel_config.pipeline_parallel_size == 1: self._modify_prev_results() self._pre_async_results = None @@ -1722,6 +1735,13 @@ def _execute_model( lora_metadata, ) + if self.is_last_rank and logits is not None: + try: + logits.block_until_ready() + logger.info(f"[PP Debug] Rank {self.rank} logits.block_until_ready() SUCCEEDED! Shape={logits.shape}") + except Exception as e: + logger.error(f"[PP Debug] Rank {self.rank} logits.block_until_ready() FAILED: {e}") + self.execute_model_state = ExecuteModelState( scheduler_output=scheduler_output, attn_metadata=attn_metadata, @@ -1756,7 +1776,7 @@ def _execute_continue_decode( self, scheduler_output: "VllmSchedulerOutput", ) -> ModelRunnerOutput | None: - if self.scheduler_config.async_scheduling and self._pre_async_results is not None: + if self.scheduler_config.async_scheduling and self._pre_async_results is not None and self.parallel_config.pipeline_parallel_size == 1: self._modify_prev_results() ( @@ -2006,7 +2026,7 @@ def _sample_from_logits( processed_bonus_logits = None if spec_decode_metadata is None: - logits = logits.astype(jnp.float32) + logger.info(f"[PP Debug] Rank {self.rank} calling sample() with logits shape={logits.shape}, dtype={logits.dtype}, do_sampling={tpu_sampling_metadata.do_sampling}") with self.maybe_forbid_compile: next_tokens, processed_logits = sample( step_rng, @@ -2014,6 +2034,7 @@ def _sample_from_logits( logits, tpu_sampling_metadata, ) + logger.info(f"[PP Debug] Rank {self.rank} sample() finished, next_tokens shape={next_tokens.shape}, type={type(next_tokens)}") else: if tpu_sampling_metadata.do_sampling: bonus_rng, rejection_rng = jax.random.split(step_rng) @@ -2169,7 +2190,10 @@ def _sample_from_logits( spec_decode_metadata) # Save the previous results - next_tokens = jax.copy_to_host_async(next_tokens) + valid_sampled_token_ids = runner_utils.host_extract_sampled_tokens( + self, spec_decode_metadata, next_tokens, + logits_indices_selector, + discard_sampled_tokens_req_indices, num_reqs) self._pre_async_results = AsyncPreResults( req_ids=req_ids, next_tokens=next_tokens, @@ -2182,13 +2206,14 @@ def _sample_from_logits( spec_decode_next_tokens=spec_decode_next_tokens, spec_decode_num_rejected_tokens=spec_decode_num_rejected_tokens, spec_decode_metadata=spec_decode_metadata, + valid_sampled_token_ids=valid_sampled_token_ids, ) # Return Model output to executor model_runner_output = ModelRunnerOutput( req_ids=req_ids, req_id_to_index=self.input_batch.req_id_to_index.copy(), - sampled_token_ids=[], # Fill in async get + sampled_token_ids=valid_sampled_token_ids, logprobs=None, prompt_logprobs_dict=prompt_logprobs_dict, pooler_output=[], @@ -2211,7 +2236,8 @@ def _sample_from_logits( req_ids_dp=req_ids_dp, padded_num_scheduled_tokens_per_dp_rank= padded_num_scheduled_tokens_per_dp_rank, - runner=self) + runner=self, + valid_sampled_token_ids=valid_sampled_token_ids) return async_model_runner_output valid_sampled_token_ids = runner_utils.host_extract_sampled_tokens( @@ -2678,7 +2704,7 @@ def _prepare_inputs(self, scheduler_output: "VllmSchedulerOutput"): num_draft_tokens[req_idx] = len(draft_token_ids) token_in_tpu_cur_input_indices_dp = {} token_in_tpu_pre_next_tokens_indices_dp = {} - if self.scheduler_config.async_scheduling and self._pre_async_results is not None: + if self.scheduler_config.async_scheduling and self._pre_async_results is not None and self.parallel_config.pipeline_parallel_size == 1: # If async previous results exists, we will prepare for the token substitution here # The actual substitution will be performed in tpu during later parts of this function. ( @@ -3107,7 +3133,7 @@ def build_shared_attn() -> SharedAttentionMetadata: shared_attention_metadata = build_shared_attn() # Async scheduling: substitute placeholder tokens for DP - if self.scheduler_config.async_scheduling and self._pre_async_results is not None: + if self.scheduler_config.async_scheduling and self._pre_async_results is not None and self.parallel_config.pipeline_parallel_size == 1: # Collect all token indices that need substitution across all DP ranks all_token_indices_to_substitute = [] all_pre_next_tokens_indices = [] @@ -3220,16 +3246,40 @@ def get_intermediate_tensor_spec(self, jax_dtype = to_jax_dtype(self.dtype) num_padded_tokens = self._get_padded_total_tokens(scheduler_output) - if self.dp_size > 1: - sharding = NamedSharding( - self.mesh, PartitionSpec(ShardingAxisName.ATTN_DATA, None)) - else: - sharding = NamedSharding(self.mesh, PartitionSpec()) + hf_conf = getattr(self.model_config, "hf_config", None) + hc_mult = getattr(hf_conf, "hc_mult", None) if hf_conf is not None else None hidden_size = self.model_config.get_hidden_size() - spec = jax.ShapeDtypeStruct(shape=(num_padded_tokens, hidden_size), - dtype=jax_dtype, - sharding=sharding) - tensor_spec = {"hidden_states": spec, "residual": spec} + + if hc_mult: + hs_shape = (num_padded_tokens, hc_mult, hidden_size) + if self.dp_size > 1: + hs_sharding = NamedSharding( + self.mesh, PartitionSpec(ShardingAxisName.ATTN_DATA, None, None)) + else: + hs_sharding = NamedSharding(self.mesh, PartitionSpec()) + else: + hs_shape = (num_padded_tokens, hidden_size) + if self.dp_size > 1: + hs_sharding = NamedSharding( + self.mesh, PartitionSpec(ShardingAxisName.ATTN_DATA, None)) + else: + hs_sharding = NamedSharding(self.mesh, PartitionSpec()) + + hs_spec = jax.ShapeDtypeStruct(shape=hs_shape, + dtype=jax_dtype, + sharding=hs_sharding) + if hc_mult: + tensor_spec = {"hidden_states": hs_spec} + else: + if self.dp_size > 1: + res_sharding = NamedSharding( + self.mesh, PartitionSpec(ShardingAxisName.ATTN_DATA, None)) + else: + res_sharding = NamedSharding(self.mesh, PartitionSpec()) + res_spec = jax.ShapeDtypeStruct(shape=(num_padded_tokens, hidden_size), + dtype=jax_dtype, + sharding=res_sharding) + tensor_spec = {"hidden_states": hs_spec, "residual": res_spec} return tensor_spec def get_uuid_for_jax_transfer(self, diff --git a/tpu_inference/runner/utils.py b/tpu_inference/runner/utils.py index d3ac2e93e8..95cd2c0d26 100644 --- a/tpu_inference/runner/utils.py +++ b/tpu_inference/runner/utils.py @@ -905,9 +905,12 @@ def host_extract_sampled_tokens( sampled_output: jnp.ndarray, logits_indices_selector: np.ndarray, discard_sampled_tokens_req_indices: list, num_reqs: int): """host retrieve the sampled tokens for the current step.""" + if hasattr(sampled_output, "_cached_valid_tokens"): + return getattr(sampled_output, "_cached_valid_tokens") next_tokens = sampled_output if spec_decode_metadata is None: - next_tokens = np.asarray(jax.device_get(next_tokens)) + if not isinstance(next_tokens, np.ndarray): + next_tokens = np.asarray(next_tokens) # Map tokens back to the pre-dp shuffling order if logits_indices_selector is not None: next_tokens = next_tokens[logits_indices_selector] @@ -923,6 +926,11 @@ def host_extract_sampled_tokens( for i in discard_sampled_tokens_req_indices: valid_sampled_token_ids[i].clear() + try: + setattr(sampled_output, "_cached_valid_tokens", valid_sampled_token_ids) + except Exception: + pass + return valid_sampled_token_ids diff --git a/tpu_inference/worker/tpu_worker.py b/tpu_inference/worker/tpu_worker.py index b6d738337a..9b737beefe 100644 --- a/tpu_inference/worker/tpu_worker.py +++ b/tpu_inference/worker/tpu_worker.py @@ -202,8 +202,15 @@ def __init__( self.devices = devices if devices is not None else [] self.device_ranks = set(device.id for device in self.devices if isinstance(device, jaxlib._jax.Device)) - self.pp_config = PPConfig(vllm_config, rank, ip, prev_worker_ip, - self.parallel_config.pipeline_parallel_size) + w_ip = os.environ.get("TPU_PP_WORKER_IP", ip) + p_ip = os.environ.get("TPU_PP_PREV_WORKER_IP", prev_worker_ip) + self.pp_config = PPConfig( + vllm_config, + rank, + w_ip, + p_ip, + self.parallel_config.pipeline_parallel_size, + ) # If model_weights is set, and we are in a distributed environment on Ray, # the driver might have overwritten `model` to its local cache path. @@ -465,6 +472,9 @@ def init_device(self, self.topology_order_id = get_device_topology_order_id( jax.local_devices(), jax.devices()) + self.is_first_rank = is_first_rank + self.is_last_rank = is_last_rank + self.model_runner = TPUModelRunner(self.vllm_config, self.devices, self.rank, is_first_rank, is_last_rank) @@ -505,8 +515,10 @@ def init_device(self, def initialize_pp_transfer_connect(self): if self.rank == 0: return - jax_parallel_state.connect(self.pp_config.prev_worker_ip, - self.rank - 1) + prev_ip = self.pp_config.prev_worker_ip + if prev_ip == "localhost" and "TPU_PP_PREV_WORKER_IP" in os.environ: + prev_ip = os.environ["TPU_PP_PREV_WORKER_IP"] + jax_parallel_state.connect(prev_ip, self.rank - 1) def determine_available_memory(self) -> int: gpu_memory_utilization = self.cache_config.gpu_memory_utilization @@ -582,21 +594,30 @@ def execute_model( # receive intermediate tensors uuid = self.model_runner.get_uuid_for_jax_transfer( scheduler_output, self.rank - 1, self.step_counter) - # TODO: this method might only works for vllm model, not sure about jax models. tensor_spec = self.model_runner.get_intermediate_tensor_spec( scheduler_output) + logger.info(f"[PP Debug] Rank {self.rank} receiving intermediate tensors with uuid={uuid} spec={tensor_spec}") intermediate_tensors_dict = get_pp_group().recv_tensor_dict( uuid, tensor_spec) + logger.info(f"[PP Debug] Rank {self.rank} successfully received intermediate tensors") intermediate_tensors = JaxIntermediateTensors( intermediate_tensors_dict) + logger.info(f"[PP Debug] Rank {self.rank} executing model runner") output = self.model_runner.execute_model(scheduler_output, intermediate_tensors) + logger.info(f"[PP Debug] Rank {self.rank} model runner execution finished, output type: {type(output)}") if isinstance(output, JaxIntermediateTensors): assert self.parallel_config.pipeline_parallel_size > 1 assert not get_pp_group().is_last_rank # send intermediate tensors + try: + for k, v in output.tensors.items(): + v.block_until_ready() + logger.info(f"[PP Debug] Rank {self.rank} forward intermediate tensors block_until_ready() SUCCEEDED!") + except Exception as e: + logger.error(f"[PP Debug] Rank {self.rank} forward intermediate tensors block_until_ready() FAILED: {e}") uuid = self.model_runner.get_uuid_for_jax_transfer( scheduler_output, self.rank, self.step_counter) get_pp_group().send_tensor_dict(uuid, output.tensors) @@ -608,7 +629,7 @@ def execute_model( # TODO(mrjunwan): Figure out if this is ok after https://github.com/vllm-project/vllm/pull/26866 if has_kv_transfer_group(): return output - return output if self.is_driver_worker else None + return output if (self.is_driver_worker or self.is_last_rank) else None def sample_tokens(self, grammar_output: GrammarOutput) -> ModelRunnerOutput: