mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 02:06:15 +00:00
add for rest of policies
This commit is contained in:
@@ -995,7 +995,14 @@ class PI0Policy(PreTrainedPolicy):
|
|||||||
|
|
||||||
# Initialize model without loading weights
|
# Initialize model without loading weights
|
||||||
# Check if dataset_stats were provided in kwargs
|
# Check if dataset_stats were provided in kwargs
|
||||||
model = cls(config, **kwargs)
|
if _transformers_available:
|
||||||
|
from transformers.modeling_utils import no_init_weights
|
||||||
|
|
||||||
|
with no_init_weights():
|
||||||
|
model = cls(config, **kwargs)
|
||||||
|
model.model.paligemma_with_expert.paligemma.tie_weights()
|
||||||
|
else:
|
||||||
|
model = cls(config, **kwargs)
|
||||||
|
|
||||||
# Now manually load and remap the state dict
|
# Now manually load and remap the state dict
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -969,6 +969,7 @@ class PI05Policy(PreTrainedPolicy):
|
|||||||
# Check if dataset_stats were provided in kwargs
|
# Check if dataset_stats were provided in kwargs
|
||||||
if _transformers_available:
|
if _transformers_available:
|
||||||
from transformers.modeling_utils import no_init_weights
|
from transformers.modeling_utils import no_init_weights
|
||||||
|
|
||||||
with no_init_weights():
|
with no_init_weights():
|
||||||
model = cls(config, **kwargs)
|
model = cls(config, **kwargs)
|
||||||
model.model.paligemma_with_expert.paligemma.tie_weights()
|
model.model.paligemma_with_expert.paligemma.tie_weights()
|
||||||
|
|||||||
@@ -895,7 +895,14 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
# Initialize model without loading weights
|
# Initialize model without loading weights
|
||||||
# Check if dataset_stats were provided in kwargs
|
# Check if dataset_stats were provided in kwargs
|
||||||
model = cls(config, **kwargs)
|
if _transformers_available:
|
||||||
|
from transformers.modeling_utils import no_init_weights
|
||||||
|
|
||||||
|
with no_init_weights():
|
||||||
|
model = cls(config, **kwargs)
|
||||||
|
model.model.paligemma_with_expert.paligemma.tie_weights()
|
||||||
|
else:
|
||||||
|
model = cls(config, **kwargs)
|
||||||
|
|
||||||
# Now manually load and remap the state dict
|
# Now manually load and remap the state dict
|
||||||
try:
|
try:
|
||||||
|
|||||||
Reference in New Issue
Block a user