fix load model

This commit is contained in:
vincentchen
2025-09-08 21:58:59 +08:00
parent 272264b516
commit ff1feb611f
4 changed files with 14 additions and 11 deletions
+1 -1
View File
@@ -207,7 +207,7 @@ class DataCollator:
self.load_processor() self.load_processor()
def load_processor(self): def load_processor(self):
processor_path = self.config["processor_path"] processor_path = self.config["pretrained_qwen_vl_path"]
action_tokenizer_path = self.config["action_tokenizer_path"] action_tokenizer_path = self.config["action_tokenizer_path"]
# Use cached processors if available # Use cached processors if available
@@ -1,6 +1,7 @@
import os import os
import torch import torch
import numpy as np import numpy as np
import glob
import torch.nn as nn import torch.nn as nn
from torchdiffeq import odeint from torchdiffeq import odeint
from dataclasses import dataclass from dataclasses import dataclass
@@ -685,7 +686,6 @@ class Qwen2_5_VLMoEForAction(Qwen2_5_VLForConditionalGeneration):
""" """
# Load model components from pretrained path # Load model components from pretrained path
model_path = os.path.join(pretrained_model_path, "model.safetensors")
config_path = os.path.join(pretrained_model_path, "config.json") config_path = os.path.join(pretrained_model_path, "config.json")
config = cls.config_class.from_pretrained(config_path) config = cls.config_class.from_pretrained(config_path)
processor = AutoProcessor.from_pretrained(pretrained_model_path, use_fast=True) processor = AutoProcessor.from_pretrained(pretrained_model_path, use_fast=True)
@@ -703,8 +703,13 @@ class Qwen2_5_VLMoEForAction(Qwen2_5_VLForConditionalGeneration):
model.resize_token_embeddings(len(processor.tokenizer)) model.resize_token_embeddings(len(processor.tokenizer))
# Load model state dict from safetensors file # Load model state dict from safetensors file
state_dict = load_file(model_path, device="cpu") safetensor_files = glob.glob(os.path.join(pretrained_model_path, "*.safetensors"))
msg = model.load_state_dict(state_dict, strict=False) state_dict = {}
for file in safetensor_files:
sd = load_file(file, device="cpu")
state_dict.update(sd)
model.load_state_dict(state_dict, strict=False)
return model return model
+1 -1
View File
@@ -108,7 +108,7 @@ class QwenVlAct_Trainer:
ValueError: If required configuration keys are missing ValueError: If required configuration keys are missing
""" """
# Validate required configuration keys # Validate required configuration keys
required_keys = ["processor_path", "qwen_vl_act_config_path", "learning_rate", "num_epoch"] required_keys = ["learning_rate", "num_epoch"]
for key in required_keys: for key in required_keys:
if key not in config: if key not in config:
raise ValueError(f"Missing required configuration key: {key}") raise ValueError(f"Missing required configuration key: {key}")
+4 -6
View File
@@ -5,9 +5,7 @@
log_name: "robotic_training" log_name: "robotic_training"
log_project: "vla_training" log_project: "vla_training"
model_type: qwen2_5 model_type: qwen2_5
processor_path: "/path/to/model/"
pretrained_qwen_vl_path: "/path/to/qwen_vl_model/" pretrained_qwen_vl_path: "/path/to/qwen_vl_model/"
qwen_vl_act_config_path: "/path/to/config.json"
action_tokenizer_path: "/path/to/fast/" action_tokenizer_path: "/path/to/fast/"
save_path: "/path/to/workspace/" save_path: "/path/to/workspace/"
@@ -52,10 +50,10 @@ agent_pos_config:
height: 1 height: 1
car_pose: 3 car_pose: 3
# Checkpoint resuming configuration # # Checkpoint resuming configuration
resume: # resume:
ckpt: "/path/to/resume_model/" # ckpt: "/path/to/resume_model/"
load_ckpt_only: true # load_ckpt_only: true
# Data configuration # Data configuration
data: data: