File size: 2,546 Bytes
7b47b6f | 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 | #!/usr/bin/env python3
import argparse
import json
from pathlib import Path
from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer, Trainer, TrainingArguments
from peft import LoraConfig, TaskType, get_peft_model
from salmonn import AudioProcessor, SalmonnConfig, SalmonnForConditionalGeneration
from salmonn.training import SalmonnCollator, SalmonnDataset
def main():
parser = argparse.ArgumentParser(description="Train SALMONN-2 with Zipformer2 and Qwen3")
parser.add_argument("--config", required=True)
parser.add_argument("--data_path", required=True)
parser.add_argument("--output_dir", required=True)
args, overrides = parser.parse_known_args()
config_data = json.loads(Path(args.config).read_text())
base_llm = config_data.pop("base_llm_name_or_path")
model_path = config_data.pop("model_name_or_path", None)
attention = config_data.pop("attn_implementation", None)
if model_path:
model = SalmonnForConditionalGeneration.from_pretrained(model_path, torch_dtype="auto")
tokenizer = AutoTokenizer.from_pretrained(model_path)
else:
qwen_config = AutoConfig.from_pretrained(base_llm)
config = SalmonnConfig(qwen_config=qwen_config.to_dict(), **config_data.pop("model"))
model = SalmonnForConditionalGeneration(config)
model.base_llm = AutoModelForCausalLM.from_pretrained(
base_llm, torch_dtype="auto", attn_implementation=attention
)
tokenizer = AutoTokenizer.from_pretrained(base_llm)
peft_values = config_data.pop("lora", None)
if peft_values:
model.base_llm = get_peft_model(
model.base_llm,
LoraConfig(task_type=TaskType.CAUSAL_LM, **peft_values),
)
if tokenizer.pad_token_id is None:
tokenizer.pad_token = tokenizer.eos_token
if model.config.inject_temporal_embedding_nl:
model.register_nl_timestamp_tokenizer(tokenizer)
training_values = config_data.pop("training")
resume = training_values.pop("resume_from_checkpoint", None)
training_values["output_dir"] = args.output_dir
training_args = TrainingArguments(**training_values)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=SalmonnDataset(args.data_path),
data_collator=SalmonnCollator(tokenizer, AudioProcessor()),
)
trainer.train(resume_from_checkpoint=resume)
trainer.save_model(args.output_dir)
tokenizer.save_pretrained(args.output_dir)
if __name__ == "__main__":
main()
|