fix: cleanup of ghost brittle stars
This commit is contained in:
parent
91cfa92834
commit
beb64ca542
1 changed files with 19 additions and 77 deletions
|
|
@ -7,7 +7,11 @@ import numpy as np
|
|||
|
||||
from brittle_star_project.evaluation.checkpoint import load_metadata, metadata_to_configs
|
||||
from brittle_star_project.evaluation.eval_env_builder import build_eval_env
|
||||
from brittle_star_project.evaluation.rollout import _get_observations, _maybe_clip_action, _target_reached
|
||||
from brittle_star_project.evaluation.rollout import (
|
||||
_get_observations,
|
||||
_maybe_clip_action,
|
||||
_target_reached,
|
||||
)
|
||||
from brittle_star_project.evaluation.video import _apply_camera_overrides, _ensure_offscreen_size
|
||||
|
||||
|
||||
|
|
@ -28,55 +32,6 @@ def _hex_to_rgba(hex_color: str, alpha: float) -> np.ndarray:
|
|||
return np.asarray([red, green, blue, float(alpha)], dtype=np.float32)
|
||||
|
||||
|
||||
def _cumulative_arc_length(points: np.ndarray) -> np.ndarray:
|
||||
if len(points) == 0:
|
||||
return np.zeros((0,), dtype=np.float32)
|
||||
|
||||
deltas = np.diff(points, axis=0)
|
||||
segment_lengths = np.linalg.norm(deltas, axis=1)
|
||||
return np.concatenate(([0.0], np.cumsum(segment_lengths))).astype(np.float32)
|
||||
|
||||
|
||||
def _sample_along_path(points: np.ndarray, count: int) -> np.ndarray:
|
||||
if len(points) == 0:
|
||||
return points
|
||||
if count <= 1 or len(points) == 1:
|
||||
return points[[0]]
|
||||
|
||||
arc = _cumulative_arc_length(points)
|
||||
total = float(arc[-1])
|
||||
if total <= 0.0:
|
||||
return points[[0] * count]
|
||||
|
||||
targets = np.linspace(0.0, total, num=count, dtype=np.float32)
|
||||
sampled = np.empty((count, points.shape[1]), dtype=np.float32)
|
||||
for idx, target in enumerate(targets):
|
||||
upper = int(np.searchsorted(arc, target, side="right"))
|
||||
lower = max(upper - 1, 0)
|
||||
if upper >= len(points):
|
||||
sampled[idx] = points[-1]
|
||||
continue
|
||||
|
||||
start = points[lower]
|
||||
end = points[upper]
|
||||
span = float(arc[upper] - arc[lower])
|
||||
if span <= 1e-8:
|
||||
sampled[idx] = start
|
||||
continue
|
||||
|
||||
weight = (float(target) - float(arc[lower])) / span
|
||||
sampled[idx] = start + weight * (end - start)
|
||||
|
||||
return sampled
|
||||
|
||||
|
||||
def _append_line(scene, mujoco, start: np.ndarray, end: np.ndarray, rgba: np.ndarray, width: float) -> None:
|
||||
geom = scene.geoms[scene.ngeom]
|
||||
mujoco.mjv_connector(geom, mujoco.mjtGeom.mjGEOM_LINE, float(width), start, end)
|
||||
geom.rgba[:] = rgba
|
||||
scene.ngeom += 1
|
||||
|
||||
|
||||
def _append_sphere(scene, mujoco, center: np.ndarray, radius: float, rgba: np.ndarray) -> None:
|
||||
geom = scene.geoms[scene.ngeom]
|
||||
mujoco.mjv_initGeom(
|
||||
|
|
@ -113,9 +68,8 @@ def main() -> None:
|
|||
parser.add_argument("model", help="Path to .flax checkpoint")
|
||||
parser.add_argument("--morphology-override", default=None)
|
||||
parser.add_argument("--output-path", required=True)
|
||||
parser.add_argument("--ghost-overlay", action="store_true")
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--max-steps", type=int, default=2000)
|
||||
parser.add_argument("--max-steps", type=int, default=5000)
|
||||
parser.add_argument("--body-name", default="BrittleStarMorphology/central_disk")
|
||||
parser.add_argument("--camera-id", type=int, default=0)
|
||||
parser.add_argument("--camera-x", type=float, default=-3.0)
|
||||
|
|
@ -127,8 +81,7 @@ def main() -> None:
|
|||
parser.add_argument("--width", type=int, default=2160)
|
||||
parser.add_argument("--height", type=int, default=960)
|
||||
parser.add_argument("--frame-stride", type=int, default=5)
|
||||
parser.add_argument("--line-color", default="#50c4ba")
|
||||
parser.add_argument("--line-width", type=float, default=2.0)
|
||||
parser.add_argument("--path-color", default="#50c4ba")
|
||||
args = parser.parse_args()
|
||||
|
||||
model_path = Path(args.model)
|
||||
|
|
@ -168,7 +121,9 @@ def main() -> None:
|
|||
|
||||
_apply_camera_overrides(
|
||||
model,
|
||||
camera_fovy={args.camera_id: float(args.camera_fovy)} if args.camera_fovy is not None else None,
|
||||
camera_fovy={args.camera_id: float(args.camera_fovy)}
|
||||
if args.camera_fovy is not None
|
||||
else None,
|
||||
camera_xyz=(
|
||||
{args.camera_id: float(args.camera_x)} if args.camera_x is not None else None,
|
||||
{args.camera_id: float(args.camera_y)} if args.camera_y is not None else None,
|
||||
|
|
@ -203,19 +158,11 @@ def main() -> None:
|
|||
raise ValueError("Need at least two rollout positions to render a path")
|
||||
|
||||
path_points = positions_arr.copy()
|
||||
path_points[:, 2] += 0.02
|
||||
path_points[:, 2] -= 0.02
|
||||
|
||||
if target_xy is not None:
|
||||
target_point = np.asarray([target_xy[0], target_xy[1], path_points[:, 2].min()], dtype=np.float32)
|
||||
else:
|
||||
target_point = None
|
||||
|
||||
ghost_count = max(1, len(path_points) // max(int(args.frame_stride), 1)) // 2
|
||||
ghost_points = _sample_along_path(path_points, ghost_count)
|
||||
|
||||
line_rgba = _hex_to_rgba(args.line_color, 0.92)
|
||||
ghost_base_rgba = _hex_to_rgba(args.line_color, 0.10)
|
||||
target_rgba = np.asarray([0.90, 0.12, 0.12, 0.95], dtype=np.float32)
|
||||
path_step = max(1, int(args.frame_stride)) * 4
|
||||
path_points_visible = path_points[::path_step]
|
||||
path_rgba = _hex_to_rgba(args.path_color, 0.92)
|
||||
|
||||
ctx = mujoco.GLContext(args.width, args.height)
|
||||
ctx.make_current()
|
||||
|
|
@ -224,7 +171,7 @@ def main() -> None:
|
|||
camera_type = _enum_value(mujoco.mjtCamera, "mjCAMERA_FIXED")
|
||||
font_scale = _enum_value(mujoco.mjtFontScale, "mjFONTSCALE_100")
|
||||
|
||||
maxgeom = int(model.ngeom + len(path_points) + len(ghost_points) + 8)
|
||||
maxgeom = int(model.ngeom + len(path_points_visible) + 8)
|
||||
scene = mujoco.MjvScene(model, maxgeom=maxgeom)
|
||||
option = mujoco.MjvOption()
|
||||
perturb = mujoco.MjvPerturb()
|
||||
|
|
@ -242,20 +189,15 @@ def main() -> None:
|
|||
|
||||
mujoco.mjv_updateScene(model, data, option, perturb, camera, catmask, scene)
|
||||
|
||||
for idx, ghost_point in enumerate(ghost_points):
|
||||
alpha = 0.10 + 0.70 * (idx / max(len(ghost_points) - 1, 1))
|
||||
ghost_color = ghost_base_rgba.copy()
|
||||
ghost_color[3] = float(alpha)
|
||||
_append_sphere(scene, mujoco, ghost_point, 0.03, ghost_color)
|
||||
|
||||
if target_point is not None:
|
||||
_append_sphere(scene, mujoco, target_point, 0.025, target_rgba)
|
||||
for idx, path_point in enumerate(path_points_visible):
|
||||
path_rgba[3] = 0.10 + 0.70 * (idx / max(len(path_points_visible) - 1, 1))
|
||||
_append_sphere(scene, mujoco, path_point, 0.03, path_rgba)
|
||||
|
||||
rgb = np.empty((args.height, args.width, 3), dtype=np.uint8)
|
||||
depth = np.empty((args.height, args.width), dtype=np.float32)
|
||||
mujoco.mjr_render(viewport, scene, context)
|
||||
mujoco.mjr_readPixels(rgb, depth, viewport, context)
|
||||
imageio.imwrite(str(output_path := Path(args.output_path)), np.flipud(rgb))
|
||||
imageio.imwrite(args.output_path, np.flipud(rgb))
|
||||
|
||||
context.free()
|
||||
finally:
|
||||
|
|
|
|||
Reference in a new issue