fixed imports

This commit is contained in:
Fabian
2022-07-14 15:19:39 +02:00
parent 855dba7fde
commit 84386fd8e4
5 changed files with 21 additions and 21 deletions
@@ -1,5 +1,5 @@
from mp_pytorch import PhaseGenerator, NormalizedRBFBasisGenerator, ZeroStartNormalizedRBFBasisGenerator
from mp_pytorch.basis_gn.rhytmic_basis import RhythmicBasisGenerator
from mp_pytorch.basis_gn import NormalizedRBFBasisGenerator, ZeroPaddingNormalizedRBFBasisGenerator
from mp_pytorch.phase_gn import PhaseGenerator
ALL_TYPES = ["rbf", "zero_rbf", "rhythmic"]
@@ -9,9 +9,10 @@ def get_basis_generator(basis_generator_type: str, phase_generator: PhaseGenerat
if basis_generator_type == "rbf":
return NormalizedRBFBasisGenerator(phase_generator, **kwargs)
elif basis_generator_type == "zero_rbf":
return ZeroStartNormalizedRBFBasisGenerator(phase_generator, **kwargs)
return ZeroPaddingNormalizedRBFBasisGenerator(phase_generator, **kwargs)
elif basis_generator_type == "rhythmic":
return RhythmicBasisGenerator(phase_generator, **kwargs)
raise NotImplementedError()
# return RhythmicBasisGenerator(phase_generator, **kwargs)
else:
raise ValueError(f"Specified basis generator type {basis_generator_type} not supported, "
f"please choose one of {ALL_TYPES}.")
@@ -1,6 +1,7 @@
from mp_pytorch import LinearPhaseGenerator, ExpDecayPhaseGenerator
from mp_pytorch.phase_gn.rhythmic_phase_generator import RhythmicPhaseGenerator
from mp_pytorch.phase_gn.smooth_phase_generator import SmoothPhaseGenerator
from mp_pytorch.phase_gn import LinearPhaseGenerator, ExpDecayPhaseGenerator
# from mp_pytorch.phase_gn.rhythmic_phase_generator import RhythmicPhaseGenerator
# from mp_pytorch.phase_gn.smooth_phase_generator import SmoothPhaseGenerator
ALL_TYPES = ["linear", "exp", "rhythmic", "smooth"]
@@ -12,9 +13,11 @@ def get_phase_generator(phase_generator_type, **kwargs):
elif phase_generator_type == "exp":
return ExpDecayPhaseGenerator(**kwargs)
elif phase_generator_type == "rhythmic":
return RhythmicPhaseGenerator(**kwargs)
raise NotImplementedError()
# return RhythmicPhaseGenerator(**kwargs)
elif phase_generator_type == "smooth":
return SmoothPhaseGenerator(**kwargs)
raise NotImplementedError()
# return SmoothPhaseGenerator(**kwargs)
else:
raise ValueError(f"Specified phase generator type {phase_generator_type} not supported, "
f"please choose one of {ALL_TYPES}.")
@@ -1,7 +1,5 @@
from mp_pytorch.basis_gn.basis_generator import BasisGenerator
from mp_pytorch.mp.dmp import DMP
from mp_pytorch.mp.idmp import IDMP
from mp_pytorch.mp.promp import ProMP
from mp_pytorch.basis_gn import BasisGenerator
from mp_pytorch.mp import ProDMP, DMP, ProMP
ALL_TYPES = ["promp", "dmp", "idmp"]
@@ -15,7 +13,7 @@ def get_trajectory_generator(
elif trajectory_generator_type == "dmp":
return DMP(basis_generator, action_dim, **kwargs)
elif trajectory_generator_type == 'idmp':
return IDMP(basis_generator, action_dim, **kwargs)
return ProDMP(basis_generator, action_dim, **kwargs)
else:
raise ValueError(f"Specified movement primitive type {trajectory_generator_type} not supported, "
f"please choose one of {ALL_TYPES}.")