File size: 6,685 Bytes
3a464db
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
import os
import logging
import random
import pytest
import torch
import torch.distributed as dist
from transformers import AutoTokenizer, AutoConfig

from vllm import distributed
from vllm.config import ParallelConfig, VllmConfig, set_current_vllm_config, get_current_vllm_config

from dinfer.model import LLaDAModelLM, LLaDAMoeModelLM
from dinfer import BlockWiseDiffusionLLM, ThresholdParallelDecoder, HierarchyDecoder
from dinfer import DiffusionLLMServing, SamplingParams
from dinfer.decoding.utils import BlockIteratorFactory
from dinfer.decoding.serving import find_continuous_ports, init_generator
from sglang.srt.layers.quantization.modelopt_quant import ModelOptFp8Config
from dinfer.model.modeling_llada2_moe_sglang import LLaDA2SGLangLM
from dinfer.decoding.diffusion_runner import ModelRunner
from dinfer import BlockIteratorFactory, KVCacheFactory, BlockDiffusionLLM
from dinfer import ThresholdParallelDecoder,CreditThresholdParallelDecoder, HierarchyDecoder, BlockWiseDiffusionLLM, IterSmoothDiffusionLLM, VicinityCacheDiffusionLLM, IterSmoothWithVicinityCacheDiffusionLLM

LLADA2_MODEL = '/mnt/infra/dulun.dl/models/LLaDA2.0-MoE-preview/LLaDA2.0-Mini-fusemoe/checkpoint-14845_fusemoe' #mini preview

sample_params = SamplingParams(threshold=0.95, cache='prefix', temperature=0., early_stop=True, cont_weight=0, prefix_look=0, 
        after_look=0, warmup_steps=0, enable_torch_compile=True, mask_id=156895, eos_id=156892, parallel_decoding='threshold', 
        use_credit=False, use_bd=True, max_length=2048)

def get_prompts(tokenizer, mask_id, device, num=1):
    prompt = "Lily can run 12 kilometers per hour for 4 hours. After that, she can run 6 kilometers per hour. How many kilometers can she run in 8 hours? "
    m = [{"role": "user", "content": prompt}, ]
    prompt = tokenizer.apply_chat_template(m, add_generation_prompt=True, tokenize=False)
    input_ids1 = torch.tensor(tokenizer(prompt)['input_ids']).to(device).unsqueeze(0)
    len1 = input_ids1.shape[1]

    if num == 2:
        prompt = "Lily can run 12 kilometers per hour for 4 hours. How many kilometers can she run in 4 hours? "
        m = [{"role": "user", "content": prompt}, ]
        prompt = tokenizer.apply_chat_template(m, add_generation_prompt=True, tokenize=False)
        input_ids2 = torch.tensor(tokenizer(prompt)['input_ids']).to(device).unsqueeze(0)
        len2 = input_ids2.shape[1]
        ret = torch.zeros(2, max(len1, len2), dtype=input_ids1.dtype)
        ret[0, 0:len1] = input_ids1
        ret[1, 0:len2] = input_ids2
    else:
        ret = input_ids1

    return ret

def get_reference_response(master_port, input_ids):
    if 'PYTEST_XDIST_WORKER' in os.environ:
        worker_num = int(os.environ['PYTEST_XDIST_WORKER'].replace('gw', ''))
        gpu_id = worker_num % torch.cuda.device_count()
    else:
        gpu_id = 0
    
    torch.cuda.set_device(gpu_id)
    device = torch.device(gpu_id)

    os.environ['MASTER_ADDR'] = 'localhost'
    os.environ['MASTER_PORT'] = str(master_port)
    from sglang.srt import distributed
    distributed.init_distributed_environment(1, 0, 'env://', 0, 'nccl')
    distributed.initialize_model_parallel(1, 1, 1, backend='nccl')
    from sglang.srt.server_args import ServerArgs
    from sglang.srt.layers.moe import initialize_moe_config
    from dinfer.model.modeling_llada2_moe_sglang import LLaDA2SGLangLM
    from dinfer.decoding.diffusion_runner import ModelRunner
    from sglang.srt.layers.dp_attention import initialize_dp_attention
    model_config = AutoConfig.from_pretrained(LLADA2_MODEL, trust_remote_code=True)


    server_args = ServerArgs(model_path=LLADA2_MODEL, enable_dp_attention=True, trust_remote_code=True, tp_size=1, dp_size = 1, pp_size = 1,
                            port=master_port+1, dist_init_addr="127.0.0.1:{}".format(master_port+2))
    try:
        from sglang.srt.server_args import set_global_server_args_for_scheduler
    except ImportError:
        pass
    else:
        set_global_server_args_for_scheduler(server_args)
    initialize_dp_attention(
        server_args=server_args,
        model_config=model_config,
    )
    initialize_moe_config(server_args)
    model = LLaDA2SGLangLM(config=model_config, expert_map_path='.').eval()
    torch.set_default_dtype(torch.bfloat16)
    model.load_weights(LLADA2_MODEL, device=device)
    initialize_moe_config(server_args)
    
    
    model = model.to(device)
    max_length = sample_params.max_length
    model = ModelRunner(model, device, server_args=server_args, max_length=max_length, enable_compile=sample_params.enable_torch_compile)

    dllm = init_generator(model, sample_params, backend='sglang', max_length=max_length)
    ref_res = dllm.generate(input_ids, gen_length=128, block_length=128).cpu()
    
    del dllm
    torch.cuda.empty_cache()
    return ref_res

def test_server_sglang():
    if 'PYTEST_XDIST_WORKER' in os.environ:
        worker_num = int(os.environ['PYTEST_XDIST_WORKER'].replace('gw', ''))
        gpu_id = worker_num % torch.cuda.device_count()
    else:
        gpu_id = 0
    
    device = torch.device(gpu_id)
    torch.cuda.set_device(gpu_id)
    print(f"[test_serving] Initializing MoE Reference on GPU {gpu_id}")
    tokenizer = AutoTokenizer.from_pretrained(LLADA2_MODEL, trust_remote_code=True, local_files_only=True)
    input_ids = get_prompts(tokenizer, mask_id=156895, device=device)   

    # obtain reference response
    port, _ = find_continuous_ports(num_ports=7)
    ref_res = get_reference_response(port, input_ids)


    print('Test sglang serving: DP == 1 and TPEP == 2')
    port, _ = find_continuous_ports(num_ports=7)
    dllm_server = DiffusionLLMServing(LLADA2_MODEL, model_type='llada2-mini', sample_params=sample_params, num_gpus=2, dp_size=1, tpep_size=2, backend='sglang',
                                    start_port=port, end_port=port+7
                                    )
    out1 = dllm_server.generate(input_ids, gen_length=128, block_length=128).cpu()
    assert torch.all(ref_res == out1)
    dllm_server.stop_serving()


    print('Test sglang serving: DP == 2 and TPEP == 2')
    port, _ = find_continuous_ports(num_ports=21)
    input_ids2 = torch.cat([input_ids, input_ids])
    dllm_server = DiffusionLLMServing(LLADA2_MODEL, model_type='llada2-mini', sample_params=sample_params, num_gpus=4, dp_size=2, tpep_size=2, backend='sglang',
                                    start_port=port, end_port=port+21
                                    )
    out2 = dllm_server.generate(input_ids2, gen_length=128, block_length=128).cpu()
    assert torch.all(ref_res[0][ref_res[0] != 156892] == out2[0][out2[0] != 156892])
    dllm_server.stop_serving()