chenbhao commited on
Commit
64ccedb
·
1 Parent(s): 9a9a9e0

Convert omni-o.pth to HF format + add from_pretrained support

Browse files

- scripts/convert_omni_o_to_hf.py: convert .pth -> config.json + model.safetensors
- src/models/vam/model.py: from_pretrained override loads encoders after meta init
- src/models/vam/model.py: __init__ skips encoder loading on meta device
- scripts/omni_o_call.py: auto-detect HF dir and load via VAM.from_pretrained
- auto_map in config.json points to modeling_omni_o.py for trust_remote_code

scripts/convert_omni_o_to_hf.py ADDED
@@ -0,0 +1,98 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os, sys, json, argparse, shutil, torch, warnings
2
+ warnings.filterwarnings('ignore')
3
+
4
+ sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'src'))
5
+ from models.vam import VAM, VAMConfig
6
+ from transformers import AutoTokenizer
7
+
8
+
9
+ def infer_config(state_dict):
10
+ hs = state_dict['model.embed_tokens.weight'].shape[1]
11
+ nl = max(int(k.split('.')[2]) for k in state_dict if k.startswith('model.layers.')) + 1
12
+ use_moe = any('expert' in k for k in state_dict if 'mlp' in k)
13
+ return dict(hidden_size=hs, num_hidden_layers=nl, use_moe=use_moe,
14
+ num_attention_heads=hs // 96, num_key_value_heads=hs // 192)
15
+
16
+
17
+ def convert(torch_path, output_dir, tokenizer_dir='checkpoint/omni/native_hf',
18
+ sensevoice_dir='checkpoint/sensevoice', siglip_dir='checkpoint/siglip',
19
+ dtype=torch.float16, device='cpu'):
20
+ os.makedirs(output_dir, exist_ok=True)
21
+
22
+ print(f'Loading checkpoint: {torch_path}')
23
+ state = torch.load(torch_path, map_location=device, weights_only=True)
24
+ cfg = infer_config(state)
25
+ print(f' hidden_size={cfg["hidden_size"]}, num_hidden_layers={cfg["num_hidden_layers"]}, use_moe={cfg["use_moe"]}')
26
+
27
+ print('Creating model...')
28
+ config = VAMConfig(**cfg)
29
+ root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
30
+ model = VAM(config,
31
+ audio_encoder_path=os.path.join(root, sensevoice_dir),
32
+ vision_model_path=os.path.join(root, siglip_dir))
33
+
34
+ missing, unexpected = model.load_state_dict(state, strict=False)
35
+ if missing:
36
+ print(f' Missing keys (expected for encoders): {len(missing)}')
37
+ if unexpected:
38
+ print(f' Unexpected keys: {len(unexpected)}')
39
+
40
+ del state
41
+ model = model.to(dtype).half()
42
+
43
+ VAM.register_for_auto_class("AutoModelForCausalLM")
44
+ VAMConfig.register_for_auto_class()
45
+
46
+ print(f'Saving to {output_dir}...')
47
+ model.save_pretrained(output_dir, safe_serialization=True)
48
+
49
+ tokenizer_path = os.path.join(root, tokenizer_dir)
50
+ if os.path.exists(tokenizer_path):
51
+ for fn in ['tokenizer.json', 'tokenizer_config.json', 'generation_config.json', 'chat_template.jinja']:
52
+ src = os.path.join(tokenizer_path, fn)
53
+ if os.path.exists(src):
54
+ shutil.copy2(src, os.path.join(output_dir, fn))
55
+ print('Tokenzier files copied')
56
+ else:
57
+ print(f'Tokenzier not found at {tokenizer_path}, skipping')
58
+
59
+ modeling_code = '''import sys, os
60
+ sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "src"))
61
+ from models.vam import VAM, VAMConfig
62
+ '''
63
+ with open(os.path.join(output_dir, 'modeling_omni_o.py'), 'w') as f:
64
+ f.write(modeling_code)
65
+
66
+ config_path = os.path.join(output_dir, 'config.json')
67
+ if os.path.exists(config_path):
68
+ with open(config_path, 'r') as f:
69
+ cf = json.load(f)
70
+ cf['auto_map'] = {
71
+ "AutoConfig": "modeling_omni_o.VAMConfig",
72
+ "AutoModelForCausalLM": "modeling_omni_o.VAM",
73
+ }
74
+ cf['model_type'] = 'omni-o'
75
+ with open(config_path, 'w') as f:
76
+ json.dump(cf, f, indent=2, ensure_ascii=False)
77
+
78
+ params = sum(p.numel() for p in model.parameters()) / 1e6
79
+ print(f'Done! Model params: {params:.2f}M')
80
+ print(f'HF model saved to: {output_dir}')
81
+ print(f' VAM.from_pretrained("{output_dir}", audio_encoder_path=..., vision_model_path=...)')
82
+ print(f' AutoModelForCausalLM.from_pretrained("{output_dir}", trust_remote_code=True)')
83
+
84
+
85
+ if __name__ == '__main__':
86
+ p = argparse.ArgumentParser(description='Convert omni-o .pth to HuggingFace format')
87
+ p.add_argument('torch_path', help='Path to .pth checkpoint')
88
+ p.add_argument('output_dir', help='Output directory for HF model')
89
+ p.add_argument('--tokenizer_dir', default='checkpoint/omni/native_hf')
90
+ p.add_argument('--sensevoice_dir', default='checkpoint/sensevoice')
91
+ p.add_argument('--siglip_dir', default='checkpoint/siglip')
92
+ p.add_argument('--dtype', default='float16')
93
+ p.add_argument('--device', default='cpu')
94
+ args = p.parse_args()
95
+
96
+ convert(args.torch_path, args.output_dir, args.tokenizer_dir,
97
+ args.sensevoice_dir, args.siglip_dir,
98
+ getattr(torch, args.dtype), args.device)
scripts/omni_o_call.py CHANGED
@@ -294,35 +294,46 @@ def init_model(args):
294
  from funasr import AutoModel
295
  M['asr'] = AutoModel(model=os.path.join(root, args.sensevoice_dir), trust_remote_code=True, device=args.device, disable_update=True)
296
 
297
- config = VAMConfig(
298
- hidden_size=args.hidden_size,
299
- num_hidden_layers=args.num_hidden_layers,
300
- num_attention_heads=args.hidden_size // 96,
301
- num_key_value_heads=args.hidden_size // 192,
302
- use_moe=args.use_moe,
303
- )
304
  ckpt_dir = os.path.join(root, args.load_from)
305
- weight = args.weight
306
- if not weight.endswith('.pth'):
307
- if args.use_moe and not weight.endswith('_moe'):
308
- weight = f'{weight}_moe.pth'
309
- else:
310
- weight = f'{weight}.pth'
311
- ckpt_path = os.path.join(ckpt_dir, weight)
312
-
313
- model = VAM(config,
314
- audio_encoder_path=os.path.join(root, args.sensevoice_dir),
315
- vision_model_path=os.path.join(root, args.siglip_dir))
316
- state = torch.load(ckpt_path, map_location='cpu', weights_only=True)
317
- missing, unexpected = model.load_state_dict(state, strict=False)
318
- if missing:
319
- print(f' Missing keys (expected for encoders): {len(missing)}')
320
- if unexpected:
321
- print(f' Unexpected keys: {len(unexpected)}')
322
- if model.audio_encoder is not None:
323
- model.audio_encoder.to(args.device)
324
- if model.vision_encoder is not None:
325
- model.vision_encoder.to(args.device)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
326
  M['model'] = model.half().eval().to(args.device)
327
 
328
  tok_dir = os.path.join(root, args.tokenizer_dir)
 
294
  from funasr import AutoModel
295
  M['asr'] = AutoModel(model=os.path.join(root, args.sensevoice_dir), trust_remote_code=True, device=args.device, disable_update=True)
296
 
 
 
 
 
 
 
 
297
  ckpt_dir = os.path.join(root, args.load_from)
298
+ is_hf = os.path.exists(os.path.join(ckpt_dir, 'config.json')) and \
299
+ (os.path.exists(os.path.join(ckpt_dir, 'model.safetensors')) or
300
+ os.path.exists(os.path.join(ckpt_dir, 'pytorch_model.bin')))
301
+
302
+ if is_hf:
303
+ model = VAM.from_pretrained(
304
+ ckpt_dir,
305
+ audio_encoder_path=os.path.join(root, args.sensevoice_dir),
306
+ vision_model_path=os.path.join(root, args.siglip_dir),
307
+ )
308
+ else:
309
+ config = VAMConfig(
310
+ hidden_size=args.hidden_size,
311
+ num_hidden_layers=args.num_hidden_layers,
312
+ num_attention_heads=args.hidden_size // 96,
313
+ num_key_value_heads=args.hidden_size // 192,
314
+ use_moe=args.use_moe,
315
+ )
316
+ weight = args.weight
317
+ if not weight.endswith('.pth'):
318
+ if args.use_moe and not weight.endswith('_moe'):
319
+ weight = f'{weight}_moe.pth'
320
+ else:
321
+ weight = f'{weight}.pth'
322
+ ckpt_path = os.path.join(ckpt_dir, weight)
323
+ model = VAM(config,
324
+ audio_encoder_path=os.path.join(root, args.sensevoice_dir),
325
+ vision_model_path=os.path.join(root, args.siglip_dir))
326
+ state = torch.load(ckpt_path, map_location='cpu', weights_only=True)
327
+ missing, unexpected = model.load_state_dict(state, strict=False)
328
+ if missing:
329
+ print(f' Missing keys (expected for encoders): {len(missing)}')
330
+ if unexpected:
331
+ print(f' Unexpected keys: {len(unexpected)}')
332
+ if model.audio_encoder is not None:
333
+ model.audio_encoder.to(args.device)
334
+ if model.vision_encoder is not None:
335
+ model.vision_encoder.to(args.device)
336
+
337
  M['model'] = model.half().eval().to(args.device)
338
 
339
  tok_dir = os.path.join(root, args.tokenizer_dir)
src/models/vam/model.py CHANGED
@@ -91,12 +91,34 @@ class VAM(LMForCausalLM):
91
  self.audio_proj = MMAudioProjector(config.audio_hidden_size, config.hidden_size)
92
  self.vision_proj = MMVisionProjector(config.image_hidden_size, config.hidden_size, target_tokens=config.image_token_len)
93
  self.audio_pad_token, self.audio_stop_token, self.audio_spk_token = config.audio_pad_token, config.audio_stop_token, config.audio_spk_token
94
- audio_encoder = SenseVoiceAudioEncoder(audio_encoder_path) if audio_encoder_path else SenseVoiceAudioEncoder()
95
- object.__setattr__(self, 'audio_encoder', audio_encoder)
96
- object.__setattr__(self, 'audio_processor', audio_encoder.processor)
97
- vision_encoder = SiglipVisionEncoder(vision_model_path) if vision_model_path else SiglipVisionEncoder()
98
- object.__setattr__(self, 'vision_encoder', vision_encoder)
99
- object.__setattr__(self, 'vision_processor', vision_encoder.processor)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
100
 
101
  @staticmethod
102
  def load_sensevoice(path):
 
91
  self.audio_proj = MMAudioProjector(config.audio_hidden_size, config.hidden_size)
92
  self.vision_proj = MMVisionProjector(config.image_hidden_size, config.hidden_size, target_tokens=config.image_token_len)
93
  self.audio_pad_token, self.audio_stop_token, self.audio_spk_token = config.audio_pad_token, config.audio_stop_token, config.audio_spk_token
94
+ meta_init = any(p.device.type == 'meta' for p in self.parameters())
95
+ if meta_init:
96
+ object.__setattr__(self, 'audio_encoder', None)
97
+ object.__setattr__(self, 'audio_processor', None)
98
+ object.__setattr__(self, 'vision_encoder', None)
99
+ object.__setattr__(self, 'vision_processor', None)
100
+ else:
101
+ audio_enc = SenseVoiceAudioEncoder(audio_encoder_path) if audio_encoder_path else SenseVoiceAudioEncoder()
102
+ object.__setattr__(self, 'audio_encoder', audio_enc)
103
+ object.__setattr__(self, 'audio_processor', audio_enc.processor)
104
+ vision_enc = SiglipVisionEncoder(vision_model_path) if vision_model_path else SiglipVisionEncoder()
105
+ object.__setattr__(self, 'vision_encoder', vision_enc)
106
+ object.__setattr__(self, 'vision_processor', vision_enc.processor)
107
+
108
+ @classmethod
109
+ def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
110
+ audio_encoder_path = kwargs.pop('audio_encoder_path', None)
111
+ vision_model_path = kwargs.pop('vision_model_path', None)
112
+ model = super().from_pretrained(pretrained_model_name_or_path, *model_args, **kwargs)
113
+ if audio_encoder_path and model.audio_encoder is None:
114
+ enc, proc = cls.load_sensevoice(audio_encoder_path)
115
+ object.__setattr__(model, 'audio_encoder', enc)
116
+ object.__setattr__(model, 'audio_processor', proc)
117
+ if vision_model_path and model.vision_encoder is None:
118
+ enc, proc = cls.load_vision(vision_model_path)
119
+ object.__setattr__(model, 'vision_encoder', enc)
120
+ object.__setattr__(model, 'vision_processor', proc)
121
+ return model
122
 
123
  @staticmethod
124
  def load_sensevoice(path):