增加log-timing

This commit is contained in:
cyy_mac
2026-07-27 15:16:58 +08:00
parent 69c9fece6c
commit 9337e0d059
2 changed files with 114 additions and 126 deletions

View File

@@ -44,8 +44,7 @@ import numpy as np
import onnxruntime as ort import onnxruntime as ort
from go1_pro_sdk import ( from go1_pro_sdk import (
MCUClient, LowCmd, MotorCmd, MotorMode, MCUClient, MotorMode, PowerProtectViolation, JOINT_NAMES,
apply_safety, PowerProtectViolation, JOINT_NAMES,
) )
@@ -63,7 +62,7 @@ except ImportError as exc:
) from exc ) from exc
DEFAULT_ONNX = HERE / "policy_30k.onnx" DEFAULT_ONNX = HERE / "policy_30k.onnx"
LOWCMD_BACKEND = "cpp_fast_lowcmd" LOWCMD_BACKEND = "cpp_checked_servo12"
SPORT_KILL_CMD = ( SPORT_KILL_CMD = (
'ssh pi@192.168.123.161 "sudo pkill -9 -f keep_sport_alive; ' 'ssh pi@192.168.123.161 "sudo pkill -9 -f keep_sport_alive; '
'sudo pkill -9 -f Legged_sport; sudo pkill -9 -f appTransit"' '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 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): def send_damping(client):
raw = get_fast_lowcmd_builder().build_encrypted_damping() raw = get_fast_lowcmd_builder().build_encrypted_damping()
client.send_raw(raw) client.send_raw(raw)
def send_hold_cmd(client, state, args): def build_servo12_raw(targets, kp, kd, state=None, position_protect_limit=None):
cmd = LowCmd() actual_q = None if state is None else motor_pos(state)
for j in range(NUM_ACTIONS): actual_tau = None if state is None else motor_tau(state)
cmd.set_motor(j, MotorCmd( pp_limit = 0.0 if position_protect_limit is None else float(position_protect_limit)
mode=MotorMode.Servo, return get_fast_lowcmd_builder().build_encrypted_servo12_checked(
q=float(DEFAULT_DOF_POS[j]), np.asarray(targets, dtype=np.float32),
dq=0.0, float(kp),
tau=0.0, float(kd),
Kp=args.kp, actual_q=actual_q,
Kd=args.kd, actual_tau=actual_tau,
)) position_protect_limit=pp_limit,
apply_safety(
cmd,
state,
power_factor=args.power_factor,
position_limit_on=True,
position_protect_limit=None,
) )
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): 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 pp_limit = args.position_protect_limit if args.position_protect_limit > 0 else None
apply_safety( send_servo12_fast(client, targets, args.kp, args.kd, state=state, position_protect_limit=pp_limit)
cmd,
state,
power_factor=args.power_factor,
position_limit_on=True,
position_protect_limit=pp_limit,
)
send_lowcmd_fast(client, cmd)
def ramp_to_default(client, args, state, logger=None, step_base=0): 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) ratio = float(i + 1) / float(ramp_steps)
target = current + ratio * (DEFAULT_DOF_POS - current) target = current + ratio * (DEFAULT_DOF_POS - current)
cmd = LowCmd() send_servo12_fast(client, target, args.kp_cal, args.kd_cal, state=state)
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)
if logger is not None and (i % log_every == 0 or i == ramp_steps - 1): 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) log_state(logger, step_base * 100000 + i, "RAMP", state, target=target)
if i % max(1, ramp_steps // 4) == 0: 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, 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( logger.log(
step, step,
mode=mode, 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, 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, joint_targets=np.zeros(NUM_ACTIONS, dtype=np.float32) if target is None else target,
state_reason=state_reason, 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 rl_step = 0
dt = 1.0 / args.rate_hz dt = 1.0 / args.rate_hz
next_t = time.perf_counter() next_t = time.perf_counter()
prev_loop_t = None
print("[INFO] R2 advances layers. L2 emergency-stops to IDLE.") print("[INFO] R2 advances layers. L2 emergency-stops to IDLE.")
print("[INFO] Ctrl+C exits with safe_stop.") print("[INFO] Ctrl+C exits with safe_stop.")
while not EXIT: 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() 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: if new_state is not None:
state = new_state state = new_state
if state is None: if state is None:
@@ -987,7 +969,10 @@ def run_deploy(args):
else: else:
ok, reason = state_ok(state, args) ok, reason = state_ok(state, args)
if ok: if ok:
policy_t0 = time.perf_counter()
action_raw = policy(onnx_input) 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) ok, reason = action_ok(action_raw, args)
if ok: if ok:
action_safe = smooth_action(clamp_action(action_raw, args), prev_action, args) action_safe = smooth_action(clamp_action(action_raw, args), prev_action, args)
@@ -1031,7 +1016,10 @@ def run_deploy(args):
else: else:
ok, reason = state_ok(state, args) ok, reason = state_ok(state, args)
if ok: if ok:
policy_t0 = time.perf_counter()
action_raw = policy(onnx_input) 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) ok, reason = action_ok(action_raw, args)
if not ok: if not ok:
print(f"\n[FAULT] RL failed: {reason}") print(f"\n[FAULT] RL failed: {reason}")
@@ -1067,6 +1055,9 @@ def run_deploy(args):
sm_state = State.HOLD sm_state = State.HOLD
print("[STATE] HOLD") print("[STATE] HOLD")
if timing is not None:
timing["work_ms"] = (time.perf_counter() - loop_t0) * 1000.0
log_state( log_state(
logger, logger,
step, step,
@@ -1079,6 +1070,7 @@ def run_deploy(args):
action_safe=action_safe, action_safe=action_safe,
target=target, target=target,
state_reason=reason, state_reason=reason,
timing=timing,
) )
if step % args.print_every == 0: 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-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-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 return parser

View File

@@ -45,8 +45,7 @@ import numpy as np
import onnxruntime as ort import onnxruntime as ort
from go1_pro_sdk import ( from go1_pro_sdk import (
MCUClient, LowCmd, MotorCmd, MotorMode, MCUClient, MotorMode, PowerProtectViolation, JOINT_NAMES,
apply_safety, PowerProtectViolation, JOINT_NAMES,
) )
@@ -64,7 +63,7 @@ except ImportError as exc:
) from exc ) from exc
DEFAULT_ONNX = HERE / "policy_robotlab_6500.onnx" DEFAULT_ONNX = HERE / "policy_robotlab_6500.onnx"
LOWCMD_BACKEND = "cpp_fast_lowcmd" LOWCMD_BACKEND = "cpp_checked_servo12"
SPORT_KILL_CMD = ( SPORT_KILL_CMD = (
'ssh pi@192.168.123.161 "sudo pkill -9 -f keep_sport_alive; ' 'ssh pi@192.168.123.161 "sudo pkill -9 -f keep_sport_alive; '
'sudo pkill -9 -f Legged_sport; sudo pkill -9 -f appTransit"' '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 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): def send_damping(client):
raw = get_fast_lowcmd_builder().build_encrypted_damping() raw = get_fast_lowcmd_builder().build_encrypted_damping()
client.send_raw(raw) client.send_raw(raw)
def send_hold_cmd(client, state, args): def build_servo12_raw(targets, kp, kd, state=None, position_protect_limit=None):
cmd = LowCmd() actual_q = None if state is None else motor_pos(state)
for j in range(NUM_ACTIONS): actual_tau = None if state is None else motor_tau(state)
cmd.set_motor(j, MotorCmd( pp_limit = 0.0 if position_protect_limit is None else float(position_protect_limit)
mode=MotorMode.Servo, return get_fast_lowcmd_builder().build_encrypted_servo12_checked(
q=float(DEFAULT_DOF_POS[j]), np.asarray(targets, dtype=np.float32),
dq=0.0, float(kp),
tau=0.0, float(kd),
Kp=args.kp, actual_q=actual_q,
Kd=args.kd, actual_tau=actual_tau,
)) position_protect_limit=pp_limit,
apply_safety(
cmd,
state,
power_factor=args.power_factor,
position_limit_on=True,
position_protect_limit=None,
) )
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): 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 pp_limit = args.position_protect_limit if args.position_protect_limit > 0 else None
apply_safety( send_servo12_fast(client, targets, args.kp, args.kd, state=state, position_protect_limit=pp_limit)
cmd,
state,
power_factor=args.power_factor,
position_limit_on=True,
position_protect_limit=pp_limit,
)
send_lowcmd_fast(client, cmd)
def ramp_to_default(client, args, state, logger=None, step_base=0): 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) ratio = float(i + 1) / float(ramp_steps)
target = current + ratio * (DEFAULT_DOF_POS - current) target = current + ratio * (DEFAULT_DOF_POS - current)
cmd = LowCmd() send_servo12_fast(client, target, args.kp_cal, args.kd_cal, state=state)
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)
if logger is not None and (i % log_every == 0 or i == ramp_steps - 1): 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) log_state(logger, step_base * 100000 + i, "RAMP", state, target=target)
if i % max(1, ramp_steps // 4) == 0: 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, 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( logger.log(
step, step,
mode=mode, 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, 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, joint_targets=np.zeros(NUM_ACTIONS, dtype=np.float32) if target is None else target,
state_reason=state_reason, 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 rl_step = 0
dt = 1.0 / args.rate_hz dt = 1.0 / args.rate_hz
next_t = time.perf_counter() next_t = time.perf_counter()
prev_loop_t = None
print("[INFO] R2 advances layers. L2 emergency-stops to IDLE.") print("[INFO] R2 advances layers. L2 emergency-stops to IDLE.")
print("[INFO] Ctrl+C exits with safe_stop.") print("[INFO] Ctrl+C exits with safe_stop.")
while not EXIT: 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() 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: if new_state is not None:
state = new_state state = new_state
if state is None: if state is None:
@@ -988,7 +970,10 @@ def run_deploy(args):
else: else:
ok, reason = state_ok(state, args) ok, reason = state_ok(state, args)
if ok: if ok:
policy_t0 = time.perf_counter()
action_raw = policy(onnx_input) 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) ok, reason = action_ok(action_raw, args)
if ok: if ok:
action_safe = smooth_action(clamp_action(action_raw, args), prev_action, args) action_safe = smooth_action(clamp_action(action_raw, args), prev_action, args)
@@ -1032,7 +1017,10 @@ def run_deploy(args):
else: else:
ok, reason = state_ok(state, args) ok, reason = state_ok(state, args)
if ok: if ok:
policy_t0 = time.perf_counter()
action_raw = policy(onnx_input) 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) ok, reason = action_ok(action_raw, args)
if not ok: if not ok:
print(f"\n[FAULT] RL failed: {reason}") print(f"\n[FAULT] RL failed: {reason}")
@@ -1068,6 +1056,9 @@ def run_deploy(args):
sm_state = State.HOLD sm_state = State.HOLD
print("[STATE] HOLD") print("[STATE] HOLD")
if timing is not None:
timing["work_ms"] = (time.perf_counter() - loop_t0) * 1000.0
log_state( log_state(
logger, logger,
step, step,
@@ -1080,6 +1071,7 @@ def run_deploy(args):
action_safe=action_safe, action_safe=action_safe,
target=target, target=target,
state_reason=reason, state_reason=reason,
timing=timing,
) )
if step % args.print_every == 0: 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-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-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 return parser