Machine Learning Trainer
Training environment to make a walk model for the hexapod generated code that will be checked
This commit is contained in:
@@ -0,0 +1,63 @@
|
||||
import numpy as np
|
||||
import math
|
||||
import torch
|
||||
import config as cfg
|
||||
import GlobalVariables as gv
|
||||
import kinematics as kin
|
||||
import DataTypes as dt
|
||||
|
||||
|
||||
class MLWalkingState:
|
||||
def __init__(self, model_path: str | None = None):
|
||||
self.model_path = model_path or "ml/checkpoints/ppo_joint_command.zip"
|
||||
self.model = None
|
||||
self._load_model()
|
||||
self.step_count = 0
|
||||
|
||||
def _load_model(self):
|
||||
try:
|
||||
from stable_baselines3 import PPO
|
||||
except ImportError:
|
||||
print("stable-baselines3 not installed: ML walking will not be available.")
|
||||
self.model = None
|
||||
return
|
||||
|
||||
try:
|
||||
self.model = PPO.load(self.model_path)
|
||||
print(f"Loaded ML walking model from {self.model_path}")
|
||||
except Exception as exc:
|
||||
print(f"Failed to load ML walking model: {exc}")
|
||||
self.model = None
|
||||
|
||||
def infer_joint_commands(self, current_rad: dt.RadArray, direction: np.ndarray) -> dt.RadArray:
|
||||
if self.model is None:
|
||||
return current_rad
|
||||
|
||||
observation = np.concatenate([current_rad.data.flatten(), direction]).astype(np.float32)
|
||||
action, _ = self.model.predict(observation, deterministic=True)
|
||||
action = np.clip(action, -1.0, 1.0).astype(np.float32)
|
||||
|
||||
new_rad = np.clip(
|
||||
current_rad.data.flatten() + action * math.radians(5.0),
|
||||
-math.pi,
|
||||
math.pi,
|
||||
).reshape((6, 3))
|
||||
return dt.RadArray(new_rad)
|
||||
|
||||
def update(self, ctx, intent, dt_step):
|
||||
if not intent.walk:
|
||||
return "idle"
|
||||
|
||||
direction = np.array([intent.move_vector.x, intent.move_vector.y, 0.0, intent.turn], dtype=np.float32)
|
||||
target_rad = self.infer_joint_commands(ctx.current_rad, direction)
|
||||
|
||||
if ctx.robotCommunication:
|
||||
ctx.robotCommunication.send_motion(target_rad)
|
||||
|
||||
if ctx.shared_sim:
|
||||
ctx.shared_sim.updatePos(target_rad)
|
||||
ctx.shared_sim.step()
|
||||
|
||||
ctx.current_rad = target_rad
|
||||
self.step_count += 1
|
||||
return None
|
||||
Reference in New Issue
Block a user