diff --git a/vllm/model_executor/models/arctic.py b/vllm/model_executor/models/arctic.py index 031b6534fb69..0c9267994b0a 100644 --- a/vllm/model_executor/models/arctic.py +++ b/vllm/model_executor/models/arctic.py @@ -16,7 +16,6 @@ get_tensor_model_parallel_world_size, tensor_model_parallel_all_reduce, ) -from vllm.logger import init_logger from vllm.model_executor.layers.activation import SiluAndMul from vllm.model_executor.layers.attention import Attention from vllm.model_executor.layers.fused_moe import fused_experts, fused_topk @@ -42,6 +41,7 @@ from .interfaces import SupportsPP, SupportsQuant from .utils import ( + AutoWeightsLoader, extract_layer_index, is_pp_missing_parameter, make_empty_intermediate_tensors_factory, @@ -49,8 +49,6 @@ maybe_prefix, ) -logger = init_logger(__name__) - class ArcticMLP(nn.Module): def __init__( @@ -384,6 +382,7 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): cache_config = vllm_config.cache_config quant_config = vllm_config.quant_config + self.config = config self.vocab_size = config.vocab_size self.embed_tokens = VocabParallelEmbedding( self.vocab_size, config.hidden_size, org_num_embeddings=self.vocab_size @@ -426,57 +425,6 @@ def forward( hidden_states = self.norm(hidden_states) return hidden_states - -class ArcticForCausalLM(nn.Module, SupportsPP, SupportsQuant): - packed_modules_mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]} - - def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): - super().__init__() - config = vllm_config.model_config.hf_config - quant_config = vllm_config.quant_config - self.config = config - self.model = ArcticModel( - vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model") - ) - self.vocab_size = config.vocab_size - self.lm_head = ParallelLMHead( - self.vocab_size, - config.hidden_size, - quant_config=quant_config, - prefix=maybe_prefix(prefix, "lm_head"), - ) - if self.config.tie_word_embeddings: - self.lm_head.weight = self.model.embed_tokens.weight - self.num_experts = config.num_local_experts - self.num_experts_per_tok = config.num_experts_per_tok - - self.logits_processor = LogitsProcessor(config.vocab_size) - self.make_empty_intermediate_tensors = ( - self.model.make_empty_intermediate_tensors - ) - - def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: - return self.model.embed_input_ids(input_ids) - - def forward( - self, - input_ids: torch.Tensor | None, - positions: torch.Tensor, - intermediate_tensors: IntermediateTensors | None = None, - inputs_embeds: torch.Tensor | None = None, - ) -> torch.Tensor | IntermediateTensors: - hidden_states = self.model( - input_ids, positions, intermediate_tensors, inputs_embeds - ) - return hidden_states - - def compute_logits( - self, - hidden_states: torch.Tensor, - ) -> torch.Tensor | None: - logits = self.logits_processor(self.lm_head, hidden_states) - return logits - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: stacked_params_mapping = [ # (param_name, shard_name, shard_id) @@ -487,41 +435,26 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: mlp_params_mapping: list[tuple[str, str, int]] = [] expert_params_mapping: list[tuple[str, str, int]] = [] - num_layers = self.config.num_hidden_layers - - for layer in range(num_layers): - mlp_params_mapping.append( - ( - f"layers.{layer}.residual_mlp.w13.weight", - f"layers.{layer}.residual_mlp.w1.weight", - 0, - ) - ) - mlp_params_mapping.append( - ( - f"layers.{layer}.residual_mlp.w13.weight", - f"layers.{layer}.residual_mlp.w3.weight", - 1, - ) - ) - if layer % 2 == 0: - # MLP layers + + for layer in range(self.config.num_hidden_layers): + is_moe_layer = (layer + 1) % self.config.moe_layer_frequency == 0 + if is_moe_layer and self.config.use_residual: mlp_params_mapping.append( ( - f"layers.{layer}.block_sparse_moe.mlp.w13.weight", - f"layers.{layer}.block_sparse_moe.mlp.w1.weight", + f"layers.{layer}.residual_mlp.w13.weight", + f"layers.{layer}.residual_mlp.w1.weight", 0, ) ) mlp_params_mapping.append( ( - f"layers.{layer}.block_sparse_moe.mlp.w13.weight", - f"layers.{layer}.block_sparse_moe.mlp.w3.weight", + f"layers.{layer}.residual_mlp.w13.weight", + f"layers.{layer}.residual_mlp.w3.weight", 1, ) ) - else: - # MoE layers + + if is_moe_layer: for expert_id in range(self.config.num_local_experts): expert_params_mapping.append( ("ws", f"experts.{expert_id}.w1.weight", expert_id) @@ -532,15 +465,25 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: expert_params_mapping.append( ("ws", f"experts.{expert_id}.w3.weight", expert_id) ) + else: + mlp_params_mapping.append( + ( + f"layers.{layer}.block_sparse_moe.mlp.w13.weight", + f"layers.{layer}.block_sparse_moe.mlp.w1.weight", + 0, + ) + ) + mlp_params_mapping.append( + ( + f"layers.{layer}.block_sparse_moe.mlp.w13.weight", + f"layers.{layer}.block_sparse_moe.mlp.w3.weight", + 1, + ) + ) params_dict = dict(self.named_parameters()) loaded_params: set[str] = set() - logger.info( - "It will take ~10 minutes loading from the 16-bit weights. " - "Alternatively, use the prequantized 8-bit weights of arctic " - "and set load-format to `sharded_state` will accelerate loading." - ) for name, loaded_weight in weights: for param_name, weight_name, shard_id in stacked_params_mapping: if weight_name not in name: @@ -585,10 +528,67 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: if is_pp_missing_parameter(name, self): continue param = params_dict[name] - weight_loader = getattr( param, "weight_loader", default_weight_loader ) weight_loader(param, loaded_weight) loaded_params.add(name) return loaded_params + + +class ArcticForCausalLM(nn.Module, SupportsPP, SupportsQuant): + packed_modules_mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]} + + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): + super().__init__() + config = vllm_config.model_config.hf_config + quant_config = vllm_config.quant_config + self.config = config + self.model = ArcticModel( + vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model") + ) + self.vocab_size = config.vocab_size + self.lm_head = ParallelLMHead( + self.vocab_size, + config.hidden_size, + quant_config=quant_config, + prefix=maybe_prefix(prefix, "lm_head"), + ) + if self.config.tie_word_embeddings: + self.lm_head.weight = self.model.embed_tokens.weight + self.num_experts = config.num_local_experts + self.num_experts_per_tok = config.num_experts_per_tok + + self.logits_processor = LogitsProcessor(config.vocab_size) + self.make_empty_intermediate_tensors = ( + self.model.make_empty_intermediate_tensors + ) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.model.embed_input_ids(input_ids) + + def forward( + self, + input_ids: torch.Tensor | None, + positions: torch.Tensor, + intermediate_tensors: IntermediateTensors | None = None, + inputs_embeds: torch.Tensor | None = None, + ) -> torch.Tensor | IntermediateTensors: + hidden_states = self.model( + input_ids, positions, intermediate_tensors, inputs_embeds + ) + return hidden_states + + def compute_logits( + self, + hidden_states: torch.Tensor, + ) -> torch.Tensor | None: + logits = self.logits_processor(self.lm_head, hidden_states) + return logits + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + loader = AutoWeightsLoader( + self, + skip_prefixes=(["lm_head."] if self.config.tie_word_embeddings else None), + ) + return loader.load_weights(weights)