Update default targets

Removed ACT since it doesn't make sense to fine-tune ACT without having it pretrained beforehand.
SmolVLA and Pi0/0.5 are much more senseful targets.
This commit is contained in:
nemo
2025-10-16 16:13:41 +02:00
parent 7f42338ae9
commit de34e876cc
+16 -12
View File
@@ -138,22 +138,26 @@ def get_default_peft_configuration(policy_type):
""" """
if policy_type == "smolvla": if policy_type == "smolvla":
return { return {
"target_modules": r"(model\.vlm_with_expert\.lm_expert\..*\.(q_proj|v_proj)|model\.action_.*|model\.state_proj.*)", "target_modules": r"(model\.vlm_with_expert\.lm_expert\..*\.(q_proj|v_proj))",
"modules_to_save": [ "modules_to_save": [
# These are inf on load otherwise # these are initialized randomly and need full-finetuning
"normalize_inputs", "state_proj",
"normalize_targets", "action_in_proj",
"unnormalize_outputs", "action_out_proj",
"action_time_mlp_in",
"action_time_mlp_out",
], ],
} }
elif policy_type == "act": elif policy_type in ("pi0", "pi05"):
return { return {
"target_modules": r"(.*_proj|.*\.action_head)", "target_modules": r".*\.gemma_expert\..*\.self_attn.(q_proj|v_proj)",
"modules_to_save": [ "modules_to_save": [
# These are inf on load otherwise # these are initialized randomly and need full-finetuning
"normalize_inputs", "state_proj",
"normalize_targets", "action_in_proj",
"unnormalize_outputs", "action_out_proj",
"action_time_mlp_in",
"action_time_mlp_out",
], ],
} }
@@ -170,7 +174,7 @@ def wrap_policy_in_peft_model(cfg, policy):
peft_config_policy = get_default_peft_configuration(cfg.policy.type) peft_config_policy = get_default_peft_configuration(cfg.policy.type)
peft_config_cli = dataclasses.asdict(cfg.peft) if cfg.peft else {} peft_config_cli = dataclasses.asdict(cfg.peft) if cfg.peft else {}
peft_config_cli['modules_to_save'] = peft_config_cli['full_training_modules'] # compatibility with PEFT peft_config_cli["modules_to_save"] = peft_config_cli["full_training_modules"] # compatibility with PEFT
peft_method_type = PeftType[peft_config_cli["method_type"].upper()] peft_method_type = PeftType[peft_config_cli["method_type"].upper()]
peft_config_cls = PEFT_TYPE_TO_CONFIG_MAPPING[peft_method_type] peft_config_cls = PEFT_TYPE_TO_CONFIG_MAPPING[peft_method_type]