GaussianMind / weights_conversion.py
ohmyzx's picture
Upload GaussianMind configuration and code
fa6f31c verified
Raw
History Blame Contribute Delete
11.8 kB
"""
Weight conversion script: merge MOSS-Audio audio components into RoboBrain.
Usage:
python weights_conversion.py
This script:
1. Loads RoboBrain weights (vision + LLM)
2. Loads MOSS-Audio weights (audio encoder + adapter + deepstack)
3. Merges them into a combined state dict for RoboBrainAudioForConditionalGeneration
4. Saves to a single safetensors file
"""
import os
import json
from collections import OrderedDict
import torch
from safetensors.torch import load_file, save_file
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
ROBOBRAIN_PATH = os.path.join(BASE_DIR, "..", "RoboBrain2.5-4B")
MOSSAUDIO_PATH = os.path.join(BASE_DIR, "..", "MOSS-Audio-4B-Instruct")
OUTPUT_PATH = BASE_DIR
def load_robobrain_weights():
print("Loading RoboBrain weights...")
path = os.path.join(ROBOBRAIN_PATH, "model.safetensors")
weights = load_file(path, device="cpu")
print(f" Loaded {len(weights)} keys from {path}")
new_weights = OrderedDict()
for key, tensor in weights.items():
if key == "lm_head.weight":
new_weights["lm_head.weight"] = tensor
elif key.startswith("model.visual."):
new_weights["model.visual." + key[len("model.visual."):]] = tensor
elif key.startswith("model.language_model."):
new_weights["model.language_model." + key[len("model.language_model."):]] = tensor
else:
print(f" Skipping unmatched key: {key}")
print(f" Mapped to {len(new_weights)} keys in target model")
return new_weights
def load_mossaudio_weights(combined: OrderedDict):
print("Loading MOSS-Audio weights...")
index_path = os.path.join(MOSSAUDIO_PATH, "model.safetensors.index.json")
with open(index_path) as f:
idx = json.load(f)
shard_files = sorted(set(idx["weight_map"].values()))
print(f" Shards: {shard_files}")
audio_count = 0
for shard_file in shard_files:
shard_path = os.path.join(MOSSAUDIO_PATH, shard_file)
weights = load_file(shard_path, device="cpu")
for key, tensor in weights.items():
if key.startswith("audio_encoder."):
combined["model.audio_encoder." + key[len("audio_encoder."):]] = tensor
audio_count += 1
elif key.startswith("audio_adapter."):
combined["model.audio_adapter." + key[len("audio_adapter."):]] = tensor
audio_count += 1
elif key.startswith("deepstack_audio_merger_list."):
combined["model.deepstack_audio_merger_list." + key[len("deepstack_audio_merger_list."):]] = tensor
audio_count += 1
print(f" Added {audio_count} audio-related keys")
return combined
def verify_weights(combined: OrderedDict):
print("\nVerifying weight compatibility...")
visual_count = sum(1 for k in combined if "visual." in k)
language_count = sum(1 for k in combined if "language_model." in k)
audio_count = sum(1 for k in combined if "audio_" in k)
lm_head_count = sum(1 for k in combined if k == "lm_head.weight")
print(f" Vision keys: {visual_count}")
print(f" Language model keys: {language_count}")
print(f" Audio keys: {audio_count}")
print(f" LM head keys: {lm_head_count}")
print(f" Total keys: {len(combined)}")
if "model.language_model.embed_tokens.weight" in combined:
embed = combined["model.language_model.embed_tokens.weight"]
print(f" Embedding shape: {list(embed.shape)}")
if "model.audio_encoder.conv1.weight" in combined:
conv = combined["model.audio_encoder.conv1.weight"]
print(f" Audio conv1 shape: {list(conv.shape)}")
if "lm_head.weight" in combined:
lm = combined["lm_head.weight"]
print(f" LM head shape: {list(lm.shape)}")
if hasattr(torch.cuda, "is_available") and torch.cuda.is_available():
gpu_mem_est = sum(t.numel() * t.element_size() for t in combined.values()) / (1024**3)
print(f"\n Estimated GPU memory: {gpu_mem_est:.2f} GB (bf16)")
def save_config_files():
print("\nCopying config files...")
robobrain_config = json.load(open(os.path.join(ROBOBRAIN_PATH, "config.json")))
mossaudio_config = json.load(open(os.path.join(MOSSAUDIO_PATH, "config.json")))
mossaudio_processor = json.load(open(os.path.join(MOSSAUDIO_PATH, "processor_config.json")))
audio_config = mossaudio_config["audio_config"]
vision_config = {
"deepstack_visual_indexes": robobrain_config["vision_config"]["deepstack_visual_indexes"],
"depth": robobrain_config["vision_config"]["depth"],
"hidden_act": robobrain_config["vision_config"]["hidden_act"],
"hidden_size": robobrain_config["vision_config"]["hidden_size"],
"in_channels": robobrain_config["vision_config"]["in_channels"],
"initializer_range": robobrain_config["vision_config"]["initializer_range"],
"intermediate_size": robobrain_config["vision_config"]["intermediate_size"],
"model_type": robobrain_config["vision_config"]["model_type"],
"num_heads": robobrain_config["vision_config"]["num_heads"],
"num_position_embeddings": robobrain_config["vision_config"]["num_position_embeddings"],
"out_hidden_size": robobrain_config["vision_config"]["out_hidden_size"],
"patch_size": robobrain_config["vision_config"]["patch_size"],
"spatial_merge_size": robobrain_config["vision_config"]["spatial_merge_size"],
"temporal_patch_size": robobrain_config["vision_config"]["temporal_patch_size"],
}
text_config = {
"attention_bias": robobrain_config["text_config"]["attention_bias"],
"attention_dropout": robobrain_config["text_config"]["attention_dropout"],
"bos_token_id": robobrain_config["text_config"]["bos_token_id"],
"dtype": robobrain_config["text_config"]["dtype"],
"eos_token_id": robobrain_config["text_config"]["eos_token_id"],
"head_dim": robobrain_config["text_config"]["head_dim"],
"hidden_act": robobrain_config["text_config"]["hidden_act"],
"hidden_size": robobrain_config["text_config"]["hidden_size"],
"initializer_range": robobrain_config["text_config"]["initializer_range"],
"intermediate_size": robobrain_config["text_config"]["intermediate_size"],
"max_position_embeddings": robobrain_config["text_config"]["max_position_embeddings"],
"model_type": robobrain_config["text_config"]["model_type"],
"num_attention_heads": robobrain_config["text_config"]["num_attention_heads"],
"num_hidden_layers": robobrain_config["text_config"]["num_hidden_layers"],
"num_key_value_heads": robobrain_config["text_config"]["num_key_value_heads"],
"rms_norm_eps": robobrain_config["text_config"]["rms_norm_eps"],
"rope_scaling": robobrain_config["text_config"]["rope_scaling"],
"rope_theta": robobrain_config["text_config"]["rope_theta"],
"tie_word_embeddings": robobrain_config["text_config"]["tie_word_embeddings"],
"use_cache": robobrain_config["text_config"]["use_cache"],
"vocab_size": robobrain_config["text_config"]["vocab_size"],
}
merged_config = {
"architectures": ["RoboBrainAudioForConditionalGeneration"],
"auto_map": {
"AutoConfig": "configuration_robobrain_audio.RoboBrainAudioConfig",
"AutoModelForCausalLM": "modeling_robobrain_audio.RoboBrainAudioForConditionalGeneration",
"AutoProcessor": "processing_robobrain_audio.RoboBrainAudioProcessor",
},
"audio_config": audio_config,
"vision_config": vision_config,
"text_config": text_config,
"audio_token_id": mossaudio_processor.get("audio_token_id", 151654),
"audio_start_token_id": mossaudio_processor.get("audio_start_id", 151669),
"audio_end_token_id": mossaudio_processor.get("audio_end_id", 151670),
"adapter_hidden_size": mossaudio_config.get("adapter_hidden_size", 8192),
"audio_deepstack_inject_layers": mossaudio_config.get("deepstack_num_inject_layers", 3),
"ignore_index": mossaudio_config.get("ignore_index", -100),
"bos_token_id": robobrain_config["text_config"]["bos_token_id"],
"eos_token_id": robobrain_config["text_config"]["eos_token_id"],
"image_token_id": robobrain_config["image_token_id"],
"video_token_id": robobrain_config["video_token_id"],
"vision_start_token_id": robobrain_config["vision_start_token_id"],
"vision_end_token_id": robobrain_config["vision_end_token_id"],
"model_type": "robobrain_audio",
"num_hidden_layers": robobrain_config["text_config"]["num_hidden_layers"],
"vocab_size": robobrain_config["text_config"]["vocab_size"],
"tie_word_embeddings": robobrain_config["text_config"]["tie_word_embeddings"],
"transformers_version": "4.57.0",
}
config_path = os.path.join(OUTPUT_PATH, "config.json")
with open(config_path, "w") as f:
json.dump(merged_config, f, indent=2)
print(f" Wrote config to {config_path}")
processor_config = {
"auto_map": {
"AutoProcessor": "processing_robobrain_audio.RoboBrainAudioProcessor"
},
"processor_class": "RoboBrainAudioProcessor",
"mel_config": mossaudio_processor.get("mel_config", {
"mel_sr": 16000,
"mel_dim": 128,
"mel_n_fft": 400,
"mel_hop_length": 160,
"mel_dtype": "bfloat16",
"use_whisper_feature_extractor": True,
}),
"enable_time_marker": mossaudio_processor.get("enable_time_marker", True),
"audio_token_id": mossaudio_processor.get("audio_token_id", 151654),
"audio_start_id": mossaudio_processor.get("audio_start_id", 151669),
"audio_end_id": mossaudio_processor.get("audio_end_id", 151670),
}
proc_path = os.path.join(OUTPUT_PATH, "processor_config.json")
with open(proc_path, "w") as f:
json.dump(processor_config, f, indent=2)
print(f" Wrote processor config to {proc_path}")
def copy_tokenizer_files():
import shutil
print("\nCopying tokenizer and preprocessing files...")
for filename in [
"tokenizer_config.json",
"tokenizer.json",
"vocab.json",
"merges.txt",
"special_tokens_map.json",
"added_tokens.json",
"preprocessor_config.json",
"video_preprocessor_config.json",
"chat_template.json",
]:
src = os.path.join(ROBOBRAIN_PATH, filename)
dst = os.path.join(OUTPUT_PATH, filename)
if os.path.exists(src):
shutil.copy2(src, dst)
print(f" Copied {filename}")
generation_config = {
"bos_token_id": 151643,
"pad_token_id": 151643,
"eos_token_id": [151645, 151643],
"do_sample": True,
"temperature": 0.7,
"top_k": 20,
"top_p": 0.8,
"repetition_penalty": 1.0,
}
gen_path = os.path.join(OUTPUT_PATH, "generation_config.json")
with open(gen_path, "w") as f:
json.dump(generation_config, f, indent=2)
print(f" Wrote generation_config.json")
def main():
print("=" * 60)
print("Merging RoboBrain + MOSS-Audio weights")
print("=" * 60)
combined = load_robobrain_weights()
combined = load_mossaudio_weights(combined)
verify_weights(combined)
print("\nSaving merged weights...")
save_path = os.path.join(OUTPUT_PATH, "model.safetensors")
save_file(combined, save_path)
print(f" Saved {len(combined)} keys to {save_path}")
save_config_files()
copy_tokenizer_files()
print("\n" + "=" * 60)
print("Done! The merged model is ready in:")
print(f" {OUTPUT_PATH}")
print("=" * 60)
if __name__ == "__main__":
main()