增加log-timing
This commit is contained in:
@@ -44,8 +44,7 @@ import numpy as np
|
||||
import onnxruntime as ort
|
||||
|
||||
from go1_pro_sdk import (
|
||||
MCUClient, LowCmd, MotorCmd, MotorMode,
|
||||
apply_safety, PowerProtectViolation, JOINT_NAMES,
|
||||
MCUClient, MotorMode, PowerProtectViolation, JOINT_NAMES,
|
||||
)
|
||||
|
||||
|
||||
@@ -63,7 +62,7 @@ except ImportError as exc:
|
||||
) from exc
|
||||
|
||||
DEFAULT_ONNX = HERE / "policy_30k.onnx"
|
||||
LOWCMD_BACKEND = "cpp_fast_lowcmd"
|
||||
LOWCMD_BACKEND = "cpp_checked_servo12"
|
||||
SPORT_KILL_CMD = (
|
||||
'ssh pi@192.168.123.161 "sudo pkill -9 -f keep_sport_alive; '
|
||||
'sudo pkill -9 -f Legged_sport; sudo pkill -9 -f appTransit"'
|
||||
@@ -393,57 +392,37 @@ def get_fast_lowcmd_builder():
|
||||
return _FAST_LOW_CMD_BUILDER
|
||||
|
||||
|
||||
def send_lowcmd_fast(client, cmd):
|
||||
raw = get_fast_lowcmd_builder().build_encrypted_lowcmd(cmd)
|
||||
client.send_raw(raw)
|
||||
|
||||
|
||||
def send_damping(client):
|
||||
raw = get_fast_lowcmd_builder().build_encrypted_damping()
|
||||
client.send_raw(raw)
|
||||
|
||||
|
||||
def send_hold_cmd(client, state, args):
|
||||
cmd = LowCmd()
|
||||
for j in range(NUM_ACTIONS):
|
||||
cmd.set_motor(j, MotorCmd(
|
||||
mode=MotorMode.Servo,
|
||||
q=float(DEFAULT_DOF_POS[j]),
|
||||
dq=0.0,
|
||||
tau=0.0,
|
||||
Kp=args.kp,
|
||||
Kd=args.kd,
|
||||
))
|
||||
apply_safety(
|
||||
cmd,
|
||||
state,
|
||||
power_factor=args.power_factor,
|
||||
position_limit_on=True,
|
||||
position_protect_limit=None,
|
||||
def build_servo12_raw(targets, kp, kd, state=None, position_protect_limit=None):
|
||||
actual_q = None if state is None else motor_pos(state)
|
||||
actual_tau = None if state is None else motor_tau(state)
|
||||
pp_limit = 0.0 if position_protect_limit is None else float(position_protect_limit)
|
||||
return get_fast_lowcmd_builder().build_encrypted_servo12_checked(
|
||||
np.asarray(targets, dtype=np.float32),
|
||||
float(kp),
|
||||
float(kd),
|
||||
actual_q=actual_q,
|
||||
actual_tau=actual_tau,
|
||||
position_protect_limit=pp_limit,
|
||||
)
|
||||
send_lowcmd_fast(client, cmd)
|
||||
|
||||
|
||||
def send_servo12_fast(client, targets, kp, kd, state=None, position_protect_limit=None):
|
||||
raw = build_servo12_raw(targets, kp, kd, state, position_protect_limit)
|
||||
client.send_raw(raw)
|
||||
|
||||
|
||||
def send_hold_cmd(client, state, args):
|
||||
send_servo12_fast(client, DEFAULT_DOF_POS, args.kp, args.kd, state=state)
|
||||
|
||||
|
||||
def send_position_cmd(client, state, targets, args):
|
||||
cmd = LowCmd()
|
||||
for j in range(NUM_ACTIONS):
|
||||
cmd.set_motor(j, MotorCmd(
|
||||
mode=MotorMode.Servo,
|
||||
q=float(targets[j]),
|
||||
dq=0.0,
|
||||
tau=0.0,
|
||||
Kp=args.kp,
|
||||
Kd=args.kd,
|
||||
))
|
||||
pp_limit = args.position_protect_limit if args.position_protect_limit > 0 else None
|
||||
apply_safety(
|
||||
cmd,
|
||||
state,
|
||||
power_factor=args.power_factor,
|
||||
position_limit_on=True,
|
||||
position_protect_limit=pp_limit,
|
||||
)
|
||||
send_lowcmd_fast(client, cmd)
|
||||
send_servo12_fast(client, targets, args.kp, args.kd, state=state, position_protect_limit=pp_limit)
|
||||
|
||||
|
||||
def ramp_to_default(client, args, state, logger=None, step_base=0):
|
||||
@@ -470,24 +449,7 @@ def ramp_to_default(client, args, state, logger=None, step_base=0):
|
||||
|
||||
ratio = float(i + 1) / float(ramp_steps)
|
||||
target = current + ratio * (DEFAULT_DOF_POS - current)
|
||||
cmd = LowCmd()
|
||||
for j in range(NUM_ACTIONS):
|
||||
cmd.set_motor(j, MotorCmd(
|
||||
mode=MotorMode.Servo,
|
||||
q=float(target[j]),
|
||||
dq=0.0,
|
||||
tau=0.0,
|
||||
Kp=args.kp_cal,
|
||||
Kd=args.kd_cal,
|
||||
))
|
||||
apply_safety(
|
||||
cmd,
|
||||
state,
|
||||
power_factor=args.power_factor,
|
||||
position_limit_on=True,
|
||||
position_protect_limit=None,
|
||||
)
|
||||
send_lowcmd_fast(client, cmd)
|
||||
send_servo12_fast(client, target, args.kp_cal, args.kd_cal, state=state)
|
||||
if logger is not None and (i % log_every == 0 or i == ramp_steps - 1):
|
||||
log_state(logger, step_base * 100000 + i, "RAMP", state, target=target)
|
||||
if i % max(1, ramp_steps // 4) == 0:
|
||||
@@ -637,7 +599,8 @@ def connect_client(args):
|
||||
|
||||
|
||||
def log_state(logger, step, mode, state, cmd=None, cmd_raw=None, obs_single=None, action_raw=None,
|
||||
action_safe=None, target=None, state_reason="ok"):
|
||||
action_safe=None, target=None, state_reason="ok", timing=None):
|
||||
timing = {} if timing is None else timing
|
||||
logger.log(
|
||||
step,
|
||||
mode=mode,
|
||||
@@ -664,6 +627,10 @@ def log_state(logger, step, mode, state, cmd=None, cmd_raw=None, obs_single=None
|
||||
action_safe=np.zeros(NUM_ACTIONS, dtype=np.float32) if action_safe is None else action_safe,
|
||||
joint_targets=np.zeros(NUM_ACTIONS, dtype=np.float32) if target is None else target,
|
||||
state_reason=state_reason,
|
||||
loop_dt_ms=float(timing.get("loop_dt_ms", 0.0)),
|
||||
recv_ms=float(timing.get("recv_ms", 0.0)),
|
||||
policy_ms=float(timing.get("policy_ms", 0.0)),
|
||||
work_ms=float(timing.get("work_ms", 0.0)),
|
||||
)
|
||||
|
||||
|
||||
@@ -885,12 +852,27 @@ def run_deploy(args):
|
||||
rl_step = 0
|
||||
dt = 1.0 / args.rate_hz
|
||||
next_t = time.perf_counter()
|
||||
prev_loop_t = None
|
||||
|
||||
print("[INFO] R2 advances layers. L2 emergency-stops to IDLE.")
|
||||
print("[INFO] Ctrl+C exits with safe_stop.")
|
||||
|
||||
while not EXIT:
|
||||
loop_t0 = time.perf_counter()
|
||||
timing = None
|
||||
if args.log_timing:
|
||||
timing = {
|
||||
"loop_dt_ms": 0.0 if prev_loop_t is None else (loop_t0 - prev_loop_t) * 1000.0,
|
||||
"recv_ms": 0.0,
|
||||
"policy_ms": 0.0,
|
||||
"work_ms": 0.0,
|
||||
}
|
||||
prev_loop_t = loop_t0
|
||||
|
||||
recv_t0 = time.perf_counter()
|
||||
new_state = client.recv_latest()
|
||||
if timing is not None:
|
||||
timing["recv_ms"] = (time.perf_counter() - recv_t0) * 1000.0
|
||||
if new_state is not None:
|
||||
state = new_state
|
||||
if state is None:
|
||||
@@ -987,7 +969,10 @@ def run_deploy(args):
|
||||
else:
|
||||
ok, reason = state_ok(state, args)
|
||||
if ok:
|
||||
policy_t0 = time.perf_counter()
|
||||
action_raw = policy(onnx_input)
|
||||
if timing is not None:
|
||||
timing["policy_ms"] += (time.perf_counter() - policy_t0) * 1000.0
|
||||
ok, reason = action_ok(action_raw, args)
|
||||
if ok:
|
||||
action_safe = smooth_action(clamp_action(action_raw, args), prev_action, args)
|
||||
@@ -1031,7 +1016,10 @@ def run_deploy(args):
|
||||
else:
|
||||
ok, reason = state_ok(state, args)
|
||||
if ok:
|
||||
policy_t0 = time.perf_counter()
|
||||
action_raw = policy(onnx_input)
|
||||
if timing is not None:
|
||||
timing["policy_ms"] += (time.perf_counter() - policy_t0) * 1000.0
|
||||
ok, reason = action_ok(action_raw, args)
|
||||
if not ok:
|
||||
print(f"\n[FAULT] RL failed: {reason}")
|
||||
@@ -1067,6 +1055,9 @@ def run_deploy(args):
|
||||
sm_state = State.HOLD
|
||||
print("[STATE] HOLD")
|
||||
|
||||
if timing is not None:
|
||||
timing["work_ms"] = (time.perf_counter() - loop_t0) * 1000.0
|
||||
|
||||
log_state(
|
||||
logger,
|
||||
step,
|
||||
@@ -1079,6 +1070,7 @@ def run_deploy(args):
|
||||
action_safe=action_safe,
|
||||
target=target,
|
||||
state_reason=reason,
|
||||
timing=timing,
|
||||
)
|
||||
|
||||
if step % args.print_every == 0:
|
||||
@@ -1195,6 +1187,8 @@ def build_arg_parser():
|
||||
|
||||
parser.add_argument("--log-dir", default=str(HERE / "logs"), help="JSONL log directory; empty disables logging")
|
||||
parser.add_argument("--log-flush-every", type=int, default=50)
|
||||
parser.add_argument("--log-timing", action="store_true",
|
||||
help="Log per-loop timing fields for deploy-mode latency diagnosis")
|
||||
return parser
|
||||
|
||||
|
||||
|
||||
@@ -45,8 +45,7 @@ import numpy as np
|
||||
import onnxruntime as ort
|
||||
|
||||
from go1_pro_sdk import (
|
||||
MCUClient, LowCmd, MotorCmd, MotorMode,
|
||||
apply_safety, PowerProtectViolation, JOINT_NAMES,
|
||||
MCUClient, MotorMode, PowerProtectViolation, JOINT_NAMES,
|
||||
)
|
||||
|
||||
|
||||
@@ -64,7 +63,7 @@ except ImportError as exc:
|
||||
) from exc
|
||||
|
||||
DEFAULT_ONNX = HERE / "policy_robotlab_6500.onnx"
|
||||
LOWCMD_BACKEND = "cpp_fast_lowcmd"
|
||||
LOWCMD_BACKEND = "cpp_checked_servo12"
|
||||
SPORT_KILL_CMD = (
|
||||
'ssh pi@192.168.123.161 "sudo pkill -9 -f keep_sport_alive; '
|
||||
'sudo pkill -9 -f Legged_sport; sudo pkill -9 -f appTransit"'
|
||||
@@ -394,57 +393,37 @@ def get_fast_lowcmd_builder():
|
||||
return _FAST_LOW_CMD_BUILDER
|
||||
|
||||
|
||||
def send_lowcmd_fast(client, cmd):
|
||||
raw = get_fast_lowcmd_builder().build_encrypted_lowcmd(cmd)
|
||||
client.send_raw(raw)
|
||||
|
||||
|
||||
def send_damping(client):
|
||||
raw = get_fast_lowcmd_builder().build_encrypted_damping()
|
||||
client.send_raw(raw)
|
||||
|
||||
|
||||
def send_hold_cmd(client, state, args):
|
||||
cmd = LowCmd()
|
||||
for j in range(NUM_ACTIONS):
|
||||
cmd.set_motor(j, MotorCmd(
|
||||
mode=MotorMode.Servo,
|
||||
q=float(DEFAULT_DOF_POS[j]),
|
||||
dq=0.0,
|
||||
tau=0.0,
|
||||
Kp=args.kp,
|
||||
Kd=args.kd,
|
||||
))
|
||||
apply_safety(
|
||||
cmd,
|
||||
state,
|
||||
power_factor=args.power_factor,
|
||||
position_limit_on=True,
|
||||
position_protect_limit=None,
|
||||
def build_servo12_raw(targets, kp, kd, state=None, position_protect_limit=None):
|
||||
actual_q = None if state is None else motor_pos(state)
|
||||
actual_tau = None if state is None else motor_tau(state)
|
||||
pp_limit = 0.0 if position_protect_limit is None else float(position_protect_limit)
|
||||
return get_fast_lowcmd_builder().build_encrypted_servo12_checked(
|
||||
np.asarray(targets, dtype=np.float32),
|
||||
float(kp),
|
||||
float(kd),
|
||||
actual_q=actual_q,
|
||||
actual_tau=actual_tau,
|
||||
position_protect_limit=pp_limit,
|
||||
)
|
||||
send_lowcmd_fast(client, cmd)
|
||||
|
||||
|
||||
def send_servo12_fast(client, targets, kp, kd, state=None, position_protect_limit=None):
|
||||
raw = build_servo12_raw(targets, kp, kd, state, position_protect_limit)
|
||||
client.send_raw(raw)
|
||||
|
||||
|
||||
def send_hold_cmd(client, state, args):
|
||||
send_servo12_fast(client, DEFAULT_DOF_POS, args.kp, args.kd, state=state)
|
||||
|
||||
|
||||
def send_position_cmd(client, state, targets, args):
|
||||
cmd = LowCmd()
|
||||
for j in range(NUM_ACTIONS):
|
||||
cmd.set_motor(j, MotorCmd(
|
||||
mode=MotorMode.Servo,
|
||||
q=float(targets[j]),
|
||||
dq=0.0,
|
||||
tau=0.0,
|
||||
Kp=args.kp,
|
||||
Kd=args.kd,
|
||||
))
|
||||
pp_limit = args.position_protect_limit if args.position_protect_limit > 0 else None
|
||||
apply_safety(
|
||||
cmd,
|
||||
state,
|
||||
power_factor=args.power_factor,
|
||||
position_limit_on=True,
|
||||
position_protect_limit=pp_limit,
|
||||
)
|
||||
send_lowcmd_fast(client, cmd)
|
||||
send_servo12_fast(client, targets, args.kp, args.kd, state=state, position_protect_limit=pp_limit)
|
||||
|
||||
|
||||
def ramp_to_default(client, args, state, logger=None, step_base=0):
|
||||
@@ -471,24 +450,7 @@ def ramp_to_default(client, args, state, logger=None, step_base=0):
|
||||
|
||||
ratio = float(i + 1) / float(ramp_steps)
|
||||
target = current + ratio * (DEFAULT_DOF_POS - current)
|
||||
cmd = LowCmd()
|
||||
for j in range(NUM_ACTIONS):
|
||||
cmd.set_motor(j, MotorCmd(
|
||||
mode=MotorMode.Servo,
|
||||
q=float(target[j]),
|
||||
dq=0.0,
|
||||
tau=0.0,
|
||||
Kp=args.kp_cal,
|
||||
Kd=args.kd_cal,
|
||||
))
|
||||
apply_safety(
|
||||
cmd,
|
||||
state,
|
||||
power_factor=args.power_factor,
|
||||
position_limit_on=True,
|
||||
position_protect_limit=None,
|
||||
)
|
||||
send_lowcmd_fast(client, cmd)
|
||||
send_servo12_fast(client, target, args.kp_cal, args.kd_cal, state=state)
|
||||
if logger is not None and (i % log_every == 0 or i == ramp_steps - 1):
|
||||
log_state(logger, step_base * 100000 + i, "RAMP", state, target=target)
|
||||
if i % max(1, ramp_steps // 4) == 0:
|
||||
@@ -638,7 +600,8 @@ def connect_client(args):
|
||||
|
||||
|
||||
def log_state(logger, step, mode, state, cmd=None, cmd_raw=None, obs_single=None, action_raw=None,
|
||||
action_safe=None, target=None, state_reason="ok"):
|
||||
action_safe=None, target=None, state_reason="ok", timing=None):
|
||||
timing = {} if timing is None else timing
|
||||
logger.log(
|
||||
step,
|
||||
mode=mode,
|
||||
@@ -665,6 +628,10 @@ def log_state(logger, step, mode, state, cmd=None, cmd_raw=None, obs_single=None
|
||||
action_safe=np.zeros(NUM_ACTIONS, dtype=np.float32) if action_safe is None else action_safe,
|
||||
joint_targets=np.zeros(NUM_ACTIONS, dtype=np.float32) if target is None else target,
|
||||
state_reason=state_reason,
|
||||
loop_dt_ms=float(timing.get("loop_dt_ms", 0.0)),
|
||||
recv_ms=float(timing.get("recv_ms", 0.0)),
|
||||
policy_ms=float(timing.get("policy_ms", 0.0)),
|
||||
work_ms=float(timing.get("work_ms", 0.0)),
|
||||
)
|
||||
|
||||
|
||||
@@ -886,12 +853,27 @@ def run_deploy(args):
|
||||
rl_step = 0
|
||||
dt = 1.0 / args.rate_hz
|
||||
next_t = time.perf_counter()
|
||||
prev_loop_t = None
|
||||
|
||||
print("[INFO] R2 advances layers. L2 emergency-stops to IDLE.")
|
||||
print("[INFO] Ctrl+C exits with safe_stop.")
|
||||
|
||||
while not EXIT:
|
||||
loop_t0 = time.perf_counter()
|
||||
timing = None
|
||||
if args.log_timing:
|
||||
timing = {
|
||||
"loop_dt_ms": 0.0 if prev_loop_t is None else (loop_t0 - prev_loop_t) * 1000.0,
|
||||
"recv_ms": 0.0,
|
||||
"policy_ms": 0.0,
|
||||
"work_ms": 0.0,
|
||||
}
|
||||
prev_loop_t = loop_t0
|
||||
|
||||
recv_t0 = time.perf_counter()
|
||||
new_state = client.recv_latest()
|
||||
if timing is not None:
|
||||
timing["recv_ms"] = (time.perf_counter() - recv_t0) * 1000.0
|
||||
if new_state is not None:
|
||||
state = new_state
|
||||
if state is None:
|
||||
@@ -988,7 +970,10 @@ def run_deploy(args):
|
||||
else:
|
||||
ok, reason = state_ok(state, args)
|
||||
if ok:
|
||||
policy_t0 = time.perf_counter()
|
||||
action_raw = policy(onnx_input)
|
||||
if timing is not None:
|
||||
timing["policy_ms"] += (time.perf_counter() - policy_t0) * 1000.0
|
||||
ok, reason = action_ok(action_raw, args)
|
||||
if ok:
|
||||
action_safe = smooth_action(clamp_action(action_raw, args), prev_action, args)
|
||||
@@ -1032,7 +1017,10 @@ def run_deploy(args):
|
||||
else:
|
||||
ok, reason = state_ok(state, args)
|
||||
if ok:
|
||||
policy_t0 = time.perf_counter()
|
||||
action_raw = policy(onnx_input)
|
||||
if timing is not None:
|
||||
timing["policy_ms"] += (time.perf_counter() - policy_t0) * 1000.0
|
||||
ok, reason = action_ok(action_raw, args)
|
||||
if not ok:
|
||||
print(f"\n[FAULT] RL failed: {reason}")
|
||||
@@ -1068,6 +1056,9 @@ def run_deploy(args):
|
||||
sm_state = State.HOLD
|
||||
print("[STATE] HOLD")
|
||||
|
||||
if timing is not None:
|
||||
timing["work_ms"] = (time.perf_counter() - loop_t0) * 1000.0
|
||||
|
||||
log_state(
|
||||
logger,
|
||||
step,
|
||||
@@ -1080,6 +1071,7 @@ def run_deploy(args):
|
||||
action_safe=action_safe,
|
||||
target=target,
|
||||
state_reason=reason,
|
||||
timing=timing,
|
||||
)
|
||||
|
||||
if step % args.print_every == 0:
|
||||
@@ -1196,6 +1188,8 @@ def build_arg_parser():
|
||||
|
||||
parser.add_argument("--log-dir", default=str(HERE / "logs"), help="JSONL log directory; empty disables logging")
|
||||
parser.add_argument("--log-flush-every", type=int, default=50)
|
||||
parser.add_argument("--log-timing", action="store_true",
|
||||
help="Log per-loop timing fields for deploy-mode latency diagnosis")
|
||||
return parser
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user