mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-29 20:49:42 +00:00
fix ux
This commit is contained in:
@@ -306,52 +306,57 @@ def _draw_dashboard(
|
||||
return dashboard
|
||||
|
||||
|
||||
def _draw_status_badge(
|
||||
def _draw_label_timeline(
|
||||
frame: np.ndarray,
|
||||
label: str,
|
||||
intervention: bool,
|
||||
episode_index: int,
|
||||
labels: np.ndarray,
|
||||
current_index: int,
|
||||
) -> None:
|
||||
del episode_index
|
||||
if label == "positive":
|
||||
text = "POSITIVE"
|
||||
color = THEME.positive
|
||||
else:
|
||||
text = "NEGATIVE"
|
||||
color = THEME.negative
|
||||
if intervention:
|
||||
text = "POSITIVE · HUMAN CORRECTION"
|
||||
|
||||
height, width = frame.shape[:2]
|
||||
horizontal_margin = max(18, round(width * 0.04))
|
||||
bar_height = max(34, round(height * 0.075))
|
||||
bar_height = max(26, round(height * 0.055))
|
||||
bottom_margin = max(14, round(height * 0.03))
|
||||
x0 = horizontal_margin
|
||||
x1 = width - horizontal_margin
|
||||
y1 = height - bottom_margin
|
||||
y0 = y1 - bar_height
|
||||
radius = max(8, bar_height // 3)
|
||||
radius = max(6, bar_height // 3)
|
||||
|
||||
_alpha_rounded_rectangle(
|
||||
frame,
|
||||
(x0, y0),
|
||||
(x1, y1),
|
||||
color,
|
||||
alpha=0.58,
|
||||
THEME.background,
|
||||
alpha=0.36,
|
||||
radius=radius,
|
||||
)
|
||||
|
||||
font_scale = max(0.5, min(0.8, width / 1100))
|
||||
text_size = cv2.getTextSize(text, cv2.FONT_HERSHEY_SIMPLEX, font_scale, 1)[0]
|
||||
text_x = x0 + (x1 - x0 - text_size[0]) // 2
|
||||
text_y = y0 + (bar_height + text_size[1]) // 2
|
||||
_draw_text(
|
||||
inset = max(3, bar_height // 8)
|
||||
inner_x0, inner_x1 = x0 + inset, x1 - inset
|
||||
inner_y0, inner_y1 = y0 + inset, y1 - inset
|
||||
timeline_width = inner_x1 - inner_x0
|
||||
count = max(len(labels), 1)
|
||||
overlay = frame.copy()
|
||||
for start, end, label in _contiguous_segments(labels):
|
||||
segment_x0 = inner_x0 + round(timeline_width * start / count)
|
||||
segment_x1 = inner_x0 + round(timeline_width * end / count)
|
||||
color = THEME.positive if label == "positive" else THEME.negative
|
||||
cv2.rectangle(
|
||||
overlay,
|
||||
(segment_x0, inner_y0),
|
||||
(max(segment_x0 + 1, segment_x1), inner_y1),
|
||||
color,
|
||||
-1,
|
||||
)
|
||||
cv2.addWeighted(overlay, 0.56, frame, 0.44, 0, dst=frame)
|
||||
|
||||
marker_x = inner_x0 + round(timeline_width * current_index / max(len(labels) - 1, 1))
|
||||
cv2.line(
|
||||
frame,
|
||||
text,
|
||||
(text_x, text_y),
|
||||
scale=font_scale,
|
||||
color=(255, 255, 255),
|
||||
thickness=1,
|
||||
(marker_x, y0 - 3),
|
||||
(marker_x, y1 + 3),
|
||||
THEME.marker,
|
||||
2,
|
||||
cv2.LINE_AA,
|
||||
)
|
||||
|
||||
|
||||
@@ -564,19 +569,17 @@ def create_advantage_video(
|
||||
end_timestamp,
|
||||
)
|
||||
written = 0
|
||||
labels = episode["advantage_label"].astype(str).to_numpy()
|
||||
try:
|
||||
for index in range(expected_frames):
|
||||
ok, frame = capture.read()
|
||||
if not ok:
|
||||
logger.warning("Video ended after %d/%d frames", index, expected_frames)
|
||||
break
|
||||
row = episode.iloc[index]
|
||||
intervention = bool(row["intervention"]) if "intervention" in episode else False
|
||||
_draw_status_badge(
|
||||
_draw_label_timeline(
|
||||
frame,
|
||||
str(row["advantage_label"]),
|
||||
intervention,
|
||||
episode_index,
|
||||
labels,
|
||||
index,
|
||||
)
|
||||
writer.write(frame)
|
||||
written += 1
|
||||
|
||||
@@ -10,7 +10,7 @@ from lerobot.scripts.lerobot_create_advantage_video import (
|
||||
_contiguous_segments,
|
||||
_create_decode_proxy,
|
||||
_draw_dashboard,
|
||||
_draw_status_badge,
|
||||
_draw_label_timeline,
|
||||
)
|
||||
|
||||
|
||||
@@ -52,12 +52,14 @@ def test_advantage_dashboard_shape():
|
||||
assert dashboard.var() > 0
|
||||
|
||||
|
||||
def test_status_badge_modifies_frame():
|
||||
def test_label_timeline_modifies_only_lower_frame_region():
|
||||
frame = np.zeros((480, 640, 3), dtype=np.uint8)
|
||||
labels = np.array(["negative", "negative", "positive", "positive"])
|
||||
|
||||
_draw_status_badge(frame, "positive", intervention=False, episode_index=12)
|
||||
_draw_label_timeline(frame, labels, current_index=2)
|
||||
|
||||
assert frame.sum() > 0
|
||||
assert frame[:350].sum() == 0
|
||||
|
||||
|
||||
@pytest.mark.skipif(shutil.which("ffmpeg") is None, reason="ffmpeg is required")
|
||||
|
||||
Reference in New Issue
Block a user