test_config.py 4.55 KB
Newer Older
1
import pytest
2
import os
3

4
from vllm.config import ModelConfig
5
from utils import models_path_prefix
6

7
MODEL_IDS_EXPECTED = [
8
9
10
    (os.path.join(models_path_prefix, "Qwen/Qwen1.5-7B"), 32768),
    (os.path.join(models_path_prefix, "mistralai/Mistral-7B-v0.1"), 4096),
    (os.path.join(models_path_prefix, "mistralai/Mistral-7B-Instruct-v0.2"), 32768),
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
]


@pytest.mark.parametrize("model_id_expected", MODEL_IDS_EXPECTED)
def test_disable_sliding_window(model_id_expected):
    model_id, expected = model_id_expected
    model_config = ModelConfig(
        model_id,
        model_id,
        tokenizer_mode="auto",
        trust_remote_code=False,
        seed=0,
        dtype="float16",
        revision=None,
        disable_sliding_window=True,
    )
    assert model_config.max_model_len == expected

29
30
31
32
33
34
35

def test_get_sliding_window():
    TEST_SLIDING_WINDOW = 4096
    # Test that the sliding window is correctly computed.
    # For Qwen1.5/Qwen2, get_sliding_window() should be None
    # when use_sliding_window is False.
    qwen2_model_config = ModelConfig(
36
37
        os.path.join(models_path_prefix, "Qwen/Qwen1.5-7B"),
        os.path.join(models_path_prefix, "Qwen/Qwen1.5-7B"),
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
        tokenizer_mode="auto",
        trust_remote_code=False,
        seed=0,
        dtype="float16",
        revision=None,
    )

    qwen2_model_config.hf_config.use_sliding_window = False
    qwen2_model_config.hf_config.sliding_window = TEST_SLIDING_WINDOW
    assert qwen2_model_config.get_sliding_window() is None

    qwen2_model_config.hf_config.use_sliding_window = True
    assert qwen2_model_config.get_sliding_window() == TEST_SLIDING_WINDOW

    mistral_model_config = ModelConfig(
53
54
        os.path.join(models_path_prefix, "mistralai/Mistral-7B-v0.1"),
        os.path.join(models_path_prefix, "mistralai/Mistral-7B-v0.1"),
55
56
57
58
59
60
61
62
63
64
        tokenizer_mode="auto",
        trust_remote_code=False,
        seed=0,
        dtype="float16",
        revision=None,
    )
    mistral_model_config.hf_config.sliding_window = None
    assert mistral_model_config.get_sliding_window() is None

    mistral_model_config.hf_config.sliding_window = TEST_SLIDING_WINDOW
65
66
67
    assert mistral_model_config.get_sliding_window() == TEST_SLIDING_WINDOW


68
def test_rope_customization():
69
    TEST_ROPE_SCALING = {"type": "dynamic", "factor": 2.0}
70
    TEST_ROPE_THETA = 16_000_000.0
71
    LONGCHAT_ROPE_SCALING = {"type": "linear", "factor": 8.0}
72
73

    llama_model_config = ModelConfig(
74
75
        os.path.join(models_path_prefix, "meta-llama/Meta-Llama-3-8B-Instruct"),
        os.path.join(models_path_prefix, "meta-llama/Meta-Llama-3-8B-Instruct"),
76
77
78
79
80
81
        tokenizer_mode="auto",
        trust_remote_code=False,
        dtype="float16",
        seed=0,
    )
    assert getattr(llama_model_config.hf_config, "rope_scaling", None) is None
82
    assert getattr(llama_model_config.hf_config, "rope_theta", None) == 500_000
83
84
85
    assert llama_model_config.max_model_len == 8192

    llama_model_config = ModelConfig(
86
87
        os.path.join(models_path_prefix, "meta-llama/Meta-Llama-3-8B-Instruct"),
        os.path.join(models_path_prefix, "meta-llama/Meta-Llama-3-8B-Instruct"),
88
89
90
91
92
        tokenizer_mode="auto",
        trust_remote_code=False,
        dtype="float16",
        seed=0,
        rope_scaling=TEST_ROPE_SCALING,
93
        rope_theta=TEST_ROPE_THETA,
94
95
96
    )
    assert getattr(llama_model_config.hf_config, "rope_scaling",
                   None) == TEST_ROPE_SCALING
97
98
    assert getattr(llama_model_config.hf_config, "rope_theta",
                   None) == TEST_ROPE_THETA
99
100
    assert llama_model_config.max_model_len == 16384

101
    longchat_model_config = ModelConfig(
102
103
        os.path.join(models_path_prefix, "lmsys/longchat-13b-16k"),
        os.path.join(models_path_prefix, "lmsys/longchat-13b-16k"),
104
105
106
107
108
109
110
111
112
113
114
115
        tokenizer_mode="auto",
        trust_remote_code=False,
        dtype="float16",
        seed=0,
    )
    # Check if LONGCHAT_ROPE_SCALING entries are in longchat_model_config
    assert all(
        longchat_model_config.hf_config.rope_scaling.get(key) == value
        for key, value in LONGCHAT_ROPE_SCALING.items())
    assert longchat_model_config.max_model_len == 16384

    longchat_model_config = ModelConfig(
116
117
        os.path.join(models_path_prefix, "lmsys/longchat-13b-16k"),
        os.path.join(models_path_prefix, "lmsys/longchat-13b-16k"),
118
119
120
121
122
123
124
125
126
        tokenizer_mode="auto",
        trust_remote_code=False,
        dtype="float16",
        seed=0,
        rope_scaling=TEST_ROPE_SCALING,
    )
    assert getattr(longchat_model_config.hf_config, "rope_scaling",
                   None) == TEST_ROPE_SCALING
    assert longchat_model_config.max_model_len == 4096