Views
No views yet
1import unittest
2
3import torch
4
5import soundfile as sf
6from qwen_omni_utils import process_mm_info
7from transformers import (
8 Qwen2_5OmniForConditionalGeneration,
9 Qwen2_5OmniPreTrainedModel,
10 Qwen2_5OmniProcessor,
11)
12
13model_id = "yujiepan/qwen2.5-omni-tiny-random"
14# model = Qwen2_5OmniModel.from_pretrained(model_id, torch_dtype="auto", device_map="auto").eval()
15# We recommend enabling flash_attention_2 for better acceleration and memory saving.
16
17Qwen2_5OmniPreTrainedModel._init_weights = unittest.mock.Mock()
18model = Qwen2_5OmniForConditionalGeneration.from_pretrained(
19 model_id,
20 torch_dtype="auto",
21 device_map="auto",
22 attn_implementation="flash_attention_2",
23).eval()
24processor = Qwen2_5OmniProcessor.from_pretrained(model_id)
25
26conversation = [
27 {
28 "role": "system",
29 "content": [
30 {"type": "text", "text": "You are Qwen, a virtual human developed by the Qwen Team, Alibaba Group, capable of perceiving auditory and visual inputs, as well as generating text and speech."}
31 ],
32 },
33 {
34 "role": "user",
35 "content": [
36 {"type": "text", "text": "Hi, can you tell me a joke?"},
37 # {"type": "audio", "audio": "https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-Audio/glass-breaking-151256.mp3"},
38 # {"type": "video", "video": "https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen2.5-Omni/draw.mp4"},
39 {"type": "image", "image": "https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-VL/assets/demo.jpeg"},
40 ],
41 },
42]
43
44# Preparation for inference
45text = processor.apply_chat_template(conversation, add_generation_prompt=True, tokenize=False)
46audios, images, videos = process_mm_info(conversation, use_audio_in_video=True)
47print('Audios:', audios)
48print('Images:', images)
49print('Videos:', videos)
50inputs = processor(text=text, audios=audios, images=images, videos=videos, return_tensors="pt", padding=True)
51inputs = inputs.to(model.device).to(model.dtype)
52
53# Inference: Generation of the output text and audio
54text_ids, audio = model.generate(
55 **inputs, use_audio_in_video=True,
56 thinker_max_new_tokens=16, talker_max_new_tokens=16,
57 temperature=0.1,
58)
59
60text = processor.batch_decode(text_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)
61print(text, '\n' * 3)
62sf.write(
63 "/tmp/output.wav",
64 audio.reshape(-1).detach().cpu().numpy(),
65 samplerate=24000,
66)1import unittest
2from pathlib import Path
3
4import torch
5
6import accelerate
7from huggingface_hub import hf_hub_download
8from transformers import (
9 AutoConfig,
10 AutoModelForCausalLM,
11 AutoTokenizer,
12 GenerationConfig,
13 Qwen2_5OmniForConditionalGeneration,
14 Qwen2_5OmniPreTrainedModel,
15 Qwen2_5OmniProcessor,
16 pipeline,
17 set_seed,
18)
19
20source_model_id = "Qwen/Qwen2.5-Omni-7B"
21save_folder = "/tmp/yujiepan/qwen2.5-omni-tiny-random"
22
23processor = Qwen2_5OmniProcessor.from_pretrained(
24 source_model_id, trust_remote_code=True,
25)
26processor.save_pretrained(save_folder)
27
28config = AutoConfig.from_pretrained(
29 source_model_id, trust_remote_code=True,
30)
31OUTPUT_DIM = 16
32config.talker_config.num_hidden_layers = 1
33config.talker_config.hidden_size = 16
34config.talker_config.embedding_size = OUTPUT_DIM
35config.talker_config.head_dim = 16
36config.talker_config.num_attention_heads = 1
37config.talker_config.num_key_value_heads = 1
38config.talker_config.intermediate_size = 32
39config.talker_config.rope_scaling['mrope_section'] = [2, 2, 4]
40assert 2 * sum(config.talker_config.rope_scaling['mrope_section']
41 ) == config.talker_config.hidden_size / config.talker_config.num_attention_heads
42
43config.thinker_config.audio_config.num_hidden_layers = 1
44config.thinker_config.audio_config.encoder_layers = 1
45config.thinker_config.audio_config.d_model = 16
46config.thinker_config.audio_config.encoder_attention_heads = 1
47config.thinker_config.audio_config.encoder_ffn_dim = 32
48config.thinker_config.audio_config.output_dim = OUTPUT_DIM
49
50config.thinker_config.text_config.num_hidden_layers = 1
51config.thinker_config.text_config.hidden_size = OUTPUT_DIM
52config.thinker_config.text_config.intermediate_size = 32
53config.thinker_config.text_config.num_attention_heads = 1
54config.thinker_config.text_config.num_key_value_heads = 1
55config.thinker_config.text_config.rope_scaling['mrope_section'] = [2, 2, 4]
56assert 2 * sum(config.thinker_config.text_config.rope_scaling['mrope_section']
57 ) == config.thinker_config.text_config.hidden_size / config.thinker_config.text_config.num_attention_heads
58
59config.thinker_config.vision_config.depth = 2
60config.thinker_config.vision_config.embed_dim = 16
61config.thinker_config.vision_config.hidden_size = 16
62config.thinker_config.vision_config.intermediate_size = 32
63config.thinker_config.vision_config.out_hidden_size = OUTPUT_DIM
64config.thinker_config.vision_config.num_heads = 1
65config.thinker_config.vision_config.fullatt_block_indexes = [1]
66
67config.token2wav_config.bigvgan_config.resblock_dilation_sizes = [[1, 3, 5]]
68config.token2wav_config.bigvgan_config.resblock_kernel_sizes = [7]
69config.token2wav_config.bigvgan_config.upsample_initial_channel = 32
70config.token2wav_config.bigvgan_config.upsample_kernel_sizes = [11, 4]
71config.token2wav_config.bigvgan_config.upsample_rates = [5, 2]
72
73config.token2wav_config.dit_config.depth = 2
74config.token2wav_config.dit_config.num_hidden_layers = 2
75config.token2wav_config.dit_config.hidden_size = 16
76config.token2wav_config.dit_config.dim = 16
77config.token2wav_config.dit_config.emb_dim = 16
78config.token2wav_config.dit_config.enc_attention_channels = 16
79config.token2wav_config.dit_config.enc_channels = [32, 32, 32]
80config.token2wav_config.dit_config.enc_dilations = [1, 3, 4]
81config.token2wav_config.dit_config.enc_kernel_sizes = [5, 3, 1]
82config.token2wav_config.dit_config.enc_dim = 16
83config.token2wav_config.dit_config.enc_emb_dim = 16
84config.token2wav_config.dit_config.enc_lin_neurons = 16
85config.token2wav_config.dit_config.head_dim = 16
86config.token2wav_config.dit_config.num_attention_heads = 1
87config.token2wav_config.dit_config.heads = 1
88config.token2wav_config.dit_config.look_ahead_layers = [1]
89config.token2wav_config.dit_config.look_backward_layers = [0]
90# avoid mismatch in vocab size because this is random model!
91config.token2wav_config.dit_config.num_embeds = config.talker_config.vocab_size
92print(config)
93
94spk_dict = torch.load(hf_hub_download(source_model_id, 'spk_dict.pt', repo_type='model'))
95for _, info in spk_dict.items():
96 info['cond'] = info['cond'][:, :config.token2wav_config.dit_config.enc_emb_dim].clone()
97torch.save(spk_dict, Path(save_folder, "spk_dict.pt"))
98
99# patch for non-affine layernorm
100Qwen2_5OmniPreTrainedModel._init_weights = unittest.mock.Mock()
101
102torch.set_default_dtype(torch.bfloat16)
103model = Qwen2_5OmniForConditionalGeneration(
104 config,
105)
106torch.set_default_dtype(torch.float32)
107model.generation_config = GenerationConfig.from_pretrained(
108 source_model_id, trust_remote_code=True,
109)
110set_seed(42)
111with torch.no_grad():
112 for name, p in sorted(model.named_parameters()):
113 torch.nn.init.normal_(p, 0, 0.5)
114 print(name, p.shape, p.dtype)
115model.save_pretrained(save_folder)