mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-29 20:49:42 +00:00
@@ -67,6 +67,7 @@ class SmolVLAConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
# Finetuning settings
|
# Finetuning settings
|
||||||
freeze_vision_encoder: bool = True
|
freeze_vision_encoder: bool = True
|
||||||
|
fine_tune_vision_encoder: bool = False
|
||||||
train_expert_only: bool = True
|
train_expert_only: bool = True
|
||||||
train_state_proj: bool = True
|
train_state_proj: bool = True
|
||||||
|
|
||||||
|
|||||||
@@ -162,6 +162,14 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
self.config = config
|
self.config = config
|
||||||
self.init_rtc_processor()
|
self.init_rtc_processor()
|
||||||
self.model = VLAFlowMatching(config, rtc_processor=self.rtc_processor)
|
self.model = VLAFlowMatching(config, rtc_processor=self.rtc_processor)
|
||||||
|
|
||||||
|
if self.config.fine_tune_vision_encoder:
|
||||||
|
self.model.vlm_with_expert.freeze_vision_encoder = False
|
||||||
|
for params in self.model.vlm_with_expert.get_vlm_model().vision_model.parameters():
|
||||||
|
params.requires_grad = True
|
||||||
|
for params in self.model.vlm_with_expert.get_vlm_model().connector.parameters():
|
||||||
|
params.requires_grad = True
|
||||||
|
|
||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user