fix load model
This commit is contained in:
@@ -207,7 +207,7 @@ class DataCollator:
|
||||
self.load_processor()
|
||||
|
||||
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"]
|
||||
|
||||
# Use cached processors if available
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
import glob
|
||||
import torch.nn as nn
|
||||
from torchdiffeq import odeint
|
||||
from dataclasses import dataclass
|
||||
@@ -685,7 +686,6 @@ class Qwen2_5_VLMoEForAction(Qwen2_5_VLForConditionalGeneration):
|
||||
"""
|
||||
|
||||
# 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 = cls.config_class.from_pretrained(config_path)
|
||||
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))
|
||||
|
||||
# Load model state dict from safetensors file
|
||||
state_dict = load_file(model_path, device="cpu")
|
||||
msg = model.load_state_dict(state_dict, strict=False)
|
||||
safetensor_files = glob.glob(os.path.join(pretrained_model_path, "*.safetensors"))
|
||||
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
|
||||
|
||||
|
||||
@@ -108,7 +108,7 @@ class QwenVlAct_Trainer:
|
||||
ValueError: If required configuration keys are missing
|
||||
"""
|
||||
# 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:
|
||||
if key not in config:
|
||||
raise ValueError(f"Missing required configuration key: {key}")
|
||||
|
||||
@@ -5,9 +5,7 @@
|
||||
log_name: "robotic_training"
|
||||
log_project: "vla_training"
|
||||
model_type: qwen2_5
|
||||
processor_path: "/path/to/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/"
|
||||
save_path: "/path/to/workspace/"
|
||||
|
||||
@@ -52,10 +50,10 @@ agent_pos_config:
|
||||
height: 1
|
||||
car_pose: 3
|
||||
|
||||
# Checkpoint resuming configuration
|
||||
resume:
|
||||
ckpt: "/path/to/resume_model/"
|
||||
load_ckpt_only: true
|
||||
# # Checkpoint resuming configuration
|
||||
# resume:
|
||||
# ckpt: "/path/to/resume_model/"
|
||||
# load_ckpt_only: true
|
||||
|
||||
# Data configuration
|
||||
data:
|
||||
|
||||
Reference in New Issue
Block a user