modelscope / modelscope/ms-swift

自己写的pt和vllm推理差异巨大 原框架的pt和vllm也有差异

Open
#7,875 1 comment 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

question stale
Dominant language
Python
Stars
15.7k
Forks
1.7k
Avg merge
1d 16h
Merged PRs (30d)
136

Description

Question Description / 问题描述

pt:

import os
import sys
import argparse
import time

# Set environment variables BEFORE importing swift modules
# This ensures that constants initialized at import time pick up these values
if 'IMAGE_MAX_TOKEN_NUM' not in os.environ:
    os.environ['IMAGE_MAX_TOKEN_NUM'] = '1024'
if 'VIDEO_MAX_TOKEN_NUM' not in os.environ:
    os.environ['VIDEO_MAX_TOKEN_NUM'] = '128'
if 'FPS_MAX_FRAMES' not in os.environ:
    os.environ['FPS_MAX_FRAMES'] = '16'

# Qwen2.5-VL specific defaults matching swift infer
if 'MAX_RATIO' not in os.environ:
    os.environ['MAX_RATIO'] = '200'
if 'FRAME_FACTOR' not in os.environ:
    os.environ['FRAME_FACTOR'] = '2'
if 'FPS' not in os.environ:
    os.environ['FPS'] = '2.0'
if 'FPS_MIN_FRAMES' not in os.environ:
    os.environ['FPS_MIN_FRAMES'] = '4'
if 'IMAGE_MIN_TOKEN_NUM' not in os.environ:
    os.environ['IMAGE_MIN_TOKEN_NUM'] = '4'
if 'SPATIAL_MERGE_SIZE' not in os.environ:
    os.environ['SPATIAL_MERGE_SIZE'] = '2'
if 'VIDEO_MIN_TOKEN_NUM' not in os.environ:
    os.environ['VIDEO_MIN_TOKEN_NUM'] = '128'

from typing import List, Optional, Dict, Any
from dataclasses import dataclass, field

import torch
from swift.llm import InferArguments, prepare_model_template, load_dataset, PtEngine, InferRequest
from swift.utils import get_logger, seed_everything, JsonlWriter

logger = get_logger()

class PtMergedInferenceRunner:
    """
    A class for merged model inference using PyTorch backend (no adapters).
    """
    def __init__(self, 
                 model_path: str,
                 val_dataset_path: str = None,
                 result_path: str = 'infer_result_pt_merged.jsonl',
                 gpu_id: str = '0',
                 max_new_tokens: int = 2048,
                 temperature: float = None,
                 top_p: float = None,
                 top_k: int = None,
                 repetition_penalty: float = None,
                 verbose: bool = True):
        
        self.infer_backend = 'pt'
        self.model_path = model_path
        self.val_dataset_path = val_dataset_path
        self.result_path = result_path
        self.max_new_tokens = max_new_tokens
        self.temperature = temperature
        self.top_p = top_p
        self.top_k = top_k
        self.repetition_penalty = repetition_penalty
        self.verbose = verbose
        self.gpu_id = gpu_id

        # 1. Configuration (Environment Variables)
        self._setup_environment()
        
        # 2. Initialize Arguments
        self.args = self._init_arguments()
        
        # Set random seed
        seed_everything(self.args.seed)
        
        self.engine = None
        self.template = None
        self.jsonl_writer = None

    def _setup_environment(self):
        """Set up necessary environment variables."""
        os.environ['CUDA_VISIBLE_DEVICES'] = self.gpu_id
        
    def _init_arguments(self) -> InferArguments:
        """Initialize InferArguments."""
        logger.info(f"Initializing InferArguments with backend: {self.infer_backend}")

        infer_args_kwargs = {
            'model': self.model_path,
            'adapters': [], # Explicitly empty for merged model
            'stream': False,
            'max_new_tokens': self.max_new_tokens,
            'temperature': self.temperature,
            'top_p': self.top_p,
            'top_k': self.top_k,
            'repetition_penalty': self.repetition_penalty,
            'load_data_args': False,
            'val_dataset': [self.val_dataset_path] if self.val_dataset_path else [],
            'infer_backend': self.infer_backend,
            'result_path': self.result_path,
            'max_length': 128000,
            'torch_dtype': 'bfloat16',
        }
        
        return InferArguments(**infer_args_kwargs)

    def load_engine(self):
        """Load the model and initialize the inference engine."""
        logger.info(f"Initializing PT Engine...")
        start_time = time.time()

        # Transformers backend
        self.model, self.template = prepare_model_template(self.args)
        
        # PtEngine initialization
        self.engine = PtEngine.from_model_template(
            self.model, 
            self.template, 
            max_batch_size=self.args.max_batch_size
        )
        self.engine.reranker_use_activation = self.args.reranker_use_activation
        
        if self.verbose:
            logger.info(f'Model structure: {self.engine.model}')
            
        end_time = time.time()
        logger.info(f"Engine Loading Time: {end_time - start_time:.4f} s")

        # Initialize JsonlWriter
        if self.args.result_path:
            logger.info(f"Results will be saved to: {self.args.result_path}")
            self.jsonl_writer = JsonlWriter(self.args.result_path)

    def run(self):
        """Execute the inference loop."""
        if not self.engine:
            self.load_engine()

        if not self.val_dataset_path:
            logger.warning("No validation dataset provided. Exiting run.")
            return

        logger.info(f"Loading validation dataset from {self.val_dataset_path}...")
        dataset_kwargs = self.args.get_dataset_kwargs()
        dataset_kwargs.pop('interleave_prob', None)

        _, val_dataset = load_dataset(
            self.args.val_dataset, 
            split_dataset_ratio=1.0, 
            shuffle=self.args.val_dataset_shuffle, 
            **dataset_kwargs
        )

        val_dataset = list(val_dataset)
        
        logger.info("Starting inference...")
        request_config = self.args.get_request_config()
        logger.info(f"Request config: {request_config}")

        results = []
        is_causal_lm = self.args.task_type == 'causal_lm'
        
        # 1. Prepare Requests
        logger.info("Preparing requests...")
        all_requests = []
        all_labels = []
        
        for i, data in enumerate(val_dataset):
            # Handle Labels
            labels = None
            if is_causal_lm:
                labels = InferRequest.remove_response(data['messages'])
            else:
                labels = data.pop('label', None)
            all_labels.append(labels)

            # Debug: Print Input IDs
            if i == 0:
                logger.info(f"Debugging Input IDs for sample {i}")
                debug_data = {'messages': data['messages']}
                for key in ['images', 'audios', 'videos', 'tools', 'objects']:
                    if key in data:
                        debug_data[key] = data[key]
                
                self.template.set_mode('pt') 
                encoded = self.template.encode(debug_data)
                input_ids = encoded.get('input_ids')
                logger.info(f"Input IDs: {input_ids}")
                logger.info(f"Input IDs Length: {len(input_ids)}")
                
            # Construct InferRequest
            infer_request_kwargs = {'messages': data['messages']}
            for key in ['images', 'audios', 'videos', 'tools', 'objects']:
                if key in data:
                    infer_request_kwargs[key] = data[key]
                    
            infer_request = InferRequest(**infer_request_kwargs)
            all_requests.append(infer_request)
            
        logger.info(f"Prepared {len(all_requests)} requests.")

        # 2. Benchmark Inference
        logger.info("Starting inference benchmark (excluding saving)...")
        start_time = time.time()
        
        try:
            infer_kwargs = {'request_config': request_config}
            # print(11111, infer_kwargs)
            # Batch Inference
            resp_list = self.engine.infer(all_requests, **infer_kwargs)
            
        except Exception as e:
            logger.error(f"Error during inference: {e}")
            import traceback
            logger.error(traceback.format_exc())
            return

        end_time = time.time()
        inference_time = end_time - start_time
        num_requests = len(all_requests)
        
        # 3. Report Statistics
        logger.info("-" * 40)
        logger.info("Inference Benchmark Report")
        logger.info("-" * 40)
        logger.info(f"Total Requests:      {num_requests}")
        logger.info(f"Total Inference Time: {inference_time:.4f} s")
        if num_requests > 0 and inference_time > 0:
            logger.info(f"Throughput:           {num_requests / inference_time:.2f} req/s")
            logger.info(f"Avg Latency:          {inference_time / num_requests * 1000:.2f} ms/req")
        logger.info("-" * 40)

        # 4. Save Results
        logger.info("Saving results...")
        for i, (data, resp) in enumerate(zip(val_dataset, resp_list)):
            if resp and len(resp.choices) > 0:
                response = resp.choices[0].message.content
            else:
                response = ""
                logger.warning(f"Empty response for sample {i}")

            data['messages'].append({'role': 'assistant', 'content': response})
            result_item = {'response': response, 'labels': all_labels[i], **data}
            results.append(result_item)

            if self.jsonl_writer:
                self.jsonl_writer.append(result_item)
                
        logger.info("Inference completed.")

def main():
    parser = argparse.ArgumentParser(description="Reproduce MS-Swift Inference with Merged Model (PyTorch)")
    
    parser.add_argument('--model_path', type=str, required=True, help="Path to the MERGED model")
    parser.add_argument('--gpu_id', type=str, default='0', help="CUDA_VISIBLE_DEVICES")
    parser.add_argument('--val_dataset', type=str, default='/data/kaelynwang/ms-swift/test_dataset.jsonl',
                        help="Path to validation dataset")
    parser.add_argument('--result_path', type=str, default='infer_result_pt_merged.jsonl',
                        help="Path to save results")
    parser.add_argument('--max_new_tokens', type=int, default=4096, help="Max new tokens to generate")
    parser.add_argument('--temperature', type=float, default=None, help="Sampling temperature")
    parser.add_argument('--top_p', type=float, default=None, help="Top-p sampling")
    parser.add_argument('--top_k', type=int, default=None, help="Top-k sampling")
    parser.add_argument('--repetition_penalty', type=float, default=None, help="Repetition penalty")

    args = parser.parse_args()

    runner = PtMergedInferenceRunner(
        model_path=args.model_path,
        val_dataset_path=args.val_dataset,
        result_path=args.result_path,
        gpu_id=args.gpu_id,
        max_new_tokens=args.max_new_tokens,
        temperature=args.temperature,
        top_p=args.top_p,
        top_k=args.top_k,
        repetition_penalty=args.repetition_penalty
    )
    
    runner.load_engine()
    runner.run()

if __name__ == "__main__":
    main()

vllm:

import os
import sys
import argparse
import time

from typing import List, Optional, Dict, Any
from dataclasses import dataclass, field

import torch
from swift.llm import InferArguments, prepare_model_template, load_dataset, InferRequest
from swift.utils import get_logger, seed_everything, JsonlWriter

logger = get_logger()

class VllmInferenceRunner:
    """
    A specialized class for running inference with vLLM using a merged model.
    No adapters are supported in this runner as the model is assumed to be fully merged.
    """
    def __init__(self, 
                 model_path: str,
                 val_dataset_path: str = None,
                 result_path: str = 'infer_result_vllm.jsonl',
                 gpu_id: str = '0',
                 max_new_tokens: int = 2048,
                 temperature: float = 0,
                 top_p: float = 1.0,
                 top_k: int = -1,
                 repetition_penalty: float = 1.0,
                 vllm_config: Dict[str, Any] = None,
                 verbose: bool = True):
        
        self.infer_backend = 'vllm'
        self.model_path = model_path
        self.val_dataset_path = val_dataset_path
        self.result_path = result_path
        self.max_new_tokens = max_new_tokens
        self.temperature = temperature
        self.top_p = top_p
        self.top_k = top_k
        self.repetition_penalty = repetition_penalty
        self.vllm_config = vllm_config or {}
        self.verbose = verbose
        self.gpu_id = gpu_id

        # 1. Configuration (Environment Variables)
        self._setup_environment()
        
        # 2. Initialize Arguments
        self.args = self._init_arguments()
        
        # Set random seed
        seed_everything(self.args.seed)
        
        self.engine = None
        self.template = None
        self.jsonl_writer = None

    def _setup_environment(self):
        """Set up necessary environment variables."""
        os.environ['CUDA_VISIBLE_DEVICES'] = self.gpu_id
        # Set defaults if not present, but allow override
        if 'IMAGE_MAX_TOKEN_NUM' not in os.environ:
            os.environ['IMAGE_MAX_TOKEN_NUM'] = '1024'
        if 'VIDEO_MAX_TOKEN_NUM' not in os.environ:
            os.environ['VIDEO_MAX_TOKEN_NUM'] = '128'
        if 'FPS_MAX_FRAMES' not in os.environ:
            os.environ['FPS_MAX_FRAMES'] = '16'
        
        # Qwen2.5-VL specific defaults matching swift infer
        if 'MAX_RATIO' not in os.environ:
            os.environ['MAX_RATIO'] = '200'
        if 'FRAME_FACTOR' not in os.environ:
            os.environ['FRAME_FACTOR'] = '2'
        if 'FPS' not in os.environ:
            os.environ['FPS'] = '2.0'
        if 'FPS_MIN_FRAMES' not in os.environ:
            os.environ['FPS_MIN_FRAMES'] = '4'
        if 'IMAGE_MIN_TOKEN_NUM' not in os.environ:
            os.environ['IMAGE_MIN_TOKEN_NUM'] = '4'
        if 'SPATIAL_MERGE_SIZE' not in os.environ:
            os.environ['SPATIAL_MERGE_SIZE'] = '2'
        if 'VIDEO_MIN_TOKEN_NUM' not in os.environ:
            os.environ['VIDEO_MIN_TOKEN_NUM'] = '128'

    def _init_arguments(self) -> InferArguments:
        """Initialize InferArguments based on configuration."""
        logger.info(f"Initializing InferArguments with backend: {self.infer_backend}")

        infer_args_kwargs = {
            'model': self.model_path,
            'stream': False,
            'max_new_tokens': self.max_new_tokens,
            'temperature': self.temperature,
            'top_p': self.top_p,
            'top_k': self.top_k,
            'repetition_penalty': self.repetition_penalty,
            'load_data_args': False,
            'val_dataset': [self.val_dataset_path] if self.val_dataset_path else [],
            'infer_backend': self.infer_backend,
            'result_path': self.result_path,
            'max_length': 128000,
            'torch_dtype': 'bfloat16',
        }
        
        # Add VLLM args
        infer_args_kwargs.update(self.vllm_config)

        return InferArguments(**infer_args_kwargs)

    def load_engine(self):
        """Load the model and initialize the vLLM engine."""
        logger.info(f"Initializing vLLM Engine from {self.model_path}...")
        start_time = time.time()

        # Lazy import to avoid top-level dependency
        try:
            from swift.llm import VllmEngine
        except ImportError:
            from swift.llm.infer.infer_engine import VllmEngine

        self.template = self.args.get_template(None)
        
        # Prepare kwargs for VllmEngine
        engine_kwargs = {
            'model_id_or_path': self.args.model,
            'model_type': self.args.model_type,
            'revision': self.args.model_revision,
            'torch_dtype': self.args.torch_dtype,
            'template': self.template,
            'reranker_use_activation': self.args.reranker_use_activation,
        }
        
        # Get VLLM specific kwargs
        engine_kwargs.update(self.args.get_vllm_engine_kwargs())
        engine_kwargs['seed'] = self.args.seed
        
        # Default safety settings for vLLM
        if 'enforce_eager' not in engine_kwargs:
             engine_kwargs['enforce_eager'] = True
        
        try:
            self.engine = VllmEngine(**engine_kwargs)
        except ValueError as e:
            if "KV cache" in str(e) and "max_model_len" in str(e):
                logger.error("\n" + "="*50)
                logger.error("VLLM OOM Error detected!")
                logger.error("The model's max_model_len is too large for the available GPU memory.")
                logger.error("Please try reducing it by adding the argument:")
                logger.error("    --vllm_max_model_len 8192")
                logger.error("="*50 + "\n")
            raise e
            
        end_time = time.time()
        logger.info(f"Engine Loading Time: {end_time - start_time:.4f} s")

        # Initialize JsonlWriter
        if self.args.result_path:
            logger.info(f"Results will be saved to: {self.args.result_path}")
            self.jsonl_writer = JsonlWriter(self.args.result_path)

    def run(self):
        """Execute the inference loop."""
        if not self.engine:
            self.load_engine()

        if not self.val_dataset_path:
            logger.warning("No validation dataset provided. Exiting run.")
            return

        logger.info(f"Loading validation dataset from {self.val_dataset_path}...")
        dataset_kwargs = self.args.get_dataset_kwargs()
        dataset_kwargs.pop('interleave_prob', None)

        _, val_dataset = load_dataset(
            self.args.val_dataset, 
            split_dataset_ratio=1.0, 
            shuffle=self.args.val_dataset_shuffle, 
            **dataset_kwargs
        )

        val_dataset = list(val_dataset)
        
        logger.info("Starting inference...")
        request_config = self.args.get_request_config()
        logger.info(f"Request config: {request_config}")

        results = []
        is_causal_lm = self.args.task_type == 'causal_lm'
        
        # 1. Prepare Requests
        logger.info("Preparing requests...")
        all_requests = []
        all_labels = []
        
        for i, data in enumerate(val_dataset):
            # Handle Labels
            labels = None
            if is_causal_lm:
                labels = InferRequest.remove_response(data['messages'])
            else:
                labels = data.pop('label', None)
            all_labels.append(labels)

            # Debug: Print Input IDs
            if i == 0:
                logger.info(f"Debugging Input IDs for sample {i}")
                debug_data = {'messages': data['messages']}
                for key in ['images', 'audios', 'videos', 'tools', 'objects']:
                    if key in data:
                        debug_data[key] = data[key]
                
                self.template.set_mode('pt') 
                encoded = self.template.encode(debug_data)
                input_ids = encoded.get('input_ids')
                logger.info(f"Input IDs: {input_ids}")
                logger.info(f"Input IDs Length: {len(input_ids)}")
                
            # Construct InferRequest
            infer_request_kwargs = {'messages': data['messages']}
            for key in ['images', 'audios', 'videos', 'tools', 'objects']:
                if key in data:
                    infer_request_kwargs[key] = data[key]
                    
            infer_request = InferRequest(**infer_request_kwargs)
            all_requests.append(infer_request)
            
        logger.info(f"Prepared {len(all_requests)} requests.")

        # 2. Benchmark Inference
        logger.info("Starting inference benchmark (excluding saving)...")
        start_time = time.time()
        
        try:
            infer_kwargs = {'request_config': request_config}
            
            # Batch Inference (No adapters needed for merged model)
            resp_list = self.engine.infer(all_requests, **infer_kwargs)
            
        except Exception as e:
            logger.error(f"Error during inference: {e}")
            import traceback
            logger.error(traceback.format_exc())
            return

        end_time = time.time()
        inference_time = end_time - start_time
        num_requests = len(all_requests)
        
        # 3. Report Statistics
        logger.info("-" * 40)
        logger.info("Inference Benchmark Report (Merged Model)")
        logger.info("-" * 40)
        logger.info(f"Total Requests:      {num_requests}")
        logger.info(f"Total Inference Time: {inference_time:.4f} s")
        if num_requests > 0 and inference_time > 0:
            logger.info(f"Throughput:           {num_requests / inference_time:.2f} req/s")
            logger.info(f"Avg Latency:          {inference_time / num_requests * 1000:.2f} ms/req")
        logger.info("-" * 40)

        # 4. Save Results (Optional)
        logger.info("Saving results...")
        for i, (data, resp) in enumerate(zip(val_dataset, resp_list)):
            # Process Response
            if resp and len(resp.choices) > 0:
                response = resp.choices[0].message.content
            else:
                response = ""
                logger.warning(f"Empty response for sample {i}")

            # Reconstruct and Save
            data['messages'].append({'role': 'assistant', 'content': response})
            result_item = {'response': response, 'labels': all_labels[i], **data}
            results.append(result_item)

            if self.jsonl_writer:
                self.jsonl_writer.append(result_item)
                
        logger.info("Inference completed.")

def main():
    parser = argparse.ArgumentParser(description="Reproduce MS-Swift Inference with Merged Model (vLLM)")
    
    # Core arguments
    parser.add_argument('--model_path', type=str, required=True, help="Path to the MERGED model")
    parser.add_argument('--gpu_id', type=str, default='0', help="CUDA_VISIBLE_DEVICES")
    parser.add_argument('--val_dataset', type=str, default='/data/kaelynwang/ms-swift/test_dataset.jsonl',
                        help="Path to validation dataset")
    parser.add_argument('--result_path', type=str, default='infer_result_vllm.jsonl',
                        help="Path to save results")
    parser.add_argument('--max_new_tokens', type=int, default=4096, help="Max new tokens to generate")
    parser.add_argument('--temperature', type=float, default=0, help="Sampling temperature")
    parser.add_argument('--top_p', type=float, default=1.0, help="Top-p sampling")
    parser.add_argument('--top_k', type=int, default=-1, help="Top-k sampling")
    parser.add_argument('--repetition_penalty', type=float, default=1.0, help="Repetition penalty")

    # VLLM Specific arguments
    parser.add_argument('--vllm_max_model_len', type=int, default=4096, help="VLLM max model len")
    parser.add_argument('--vllm_gpu_memory_utilization', type=float, default=0.95, help="VLLM GPU memory utilization")
    parser.add_argument('--vllm_tensor_parallel_size', type=int, default=1, help="VLLM tensor parallel size")
    parser.add_argument('--vllm_enforce_eager', action='store_true', default=True, help="VLLM enforce eager mode (Default: True)")
    parser.add_argument('--vllm_max_num_seqs', type=int, default=256, help="VLLM max num seqs")

    args = parser.parse_args()

    # Construct VLLM config
    vllm_config = {
        'vllm_max_model_len': args.vllm_max_model_len,
        'vllm_gpu_memory_utilization': args.vllm_gpu_memory_utilization,
        'vllm_tensor_parallel_size': args.vllm_tensor_parallel_size,
        'vllm_enforce_eager': args.vllm_enforce_eager,
        'vllm_max_num_seqs': args.vllm_max_num_seqs,
    }

    runner = VllmInferenceRunner(
        model_path=args.model_path,
        val_dataset_path=args.val_dataset,
        result_path=args.result_path,
        gpu_id=args.gpu_id,
        max_new_tokens=args.max_new_tokens,
        temperature=args.temperature,
        top_p=args.top_p,
        top_k=args.top_k,
        repetition_penalty=args.repetition_penalty,
        vllm_config=vllm_config
    )
    
    runner.load_engine()
    runner.run()

if __name__ == "__main__":
    main()

pt和vllm 本身受推理框架的影响 一定不会100%一致的
在输入一致的前提下有没有什么更好的优化方法

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Compare the pasted PtMergedInferenceRunner and VllmInferenceRunner implementations, starting with their backend-specific initialization, environment variables, and request configurations. Reproduce both runs with the same merged model and validation dataset, then compare the logged input IDs, generation settings, and outputs to identify the source of the discrepancy; no repository test or target file is named in the issue.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
backend, machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Quiet
Clarity
Needs clarification
Newbie friendliness
38/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.