diff --git a/src/unilab/tasks/motion_tracking/common/motion_loader.py b/src/unilab/tasks/motion_tracking/common/motion_loader.py index 663f94dca..653bd319d 100644 --- a/src/unilab/tasks/motion_tracking/common/motion_loader.py +++ b/src/unilab/tasks/motion_tracking/common/motion_loader.py @@ -510,8 +510,13 @@ def _sample_adaptive(self, env_ids: np.ndarray) -> np.ndarray: padded = np.pad(sampling_probs, (0, self.adaptive_kernel_size - 1), mode="edge") sampling_probs = np.convolve(padded, self.kernel, mode="valid") - # Normalize to probabilities - sampling_probs = sampling_probs / sampling_probs.sum() + # With no failure statistics yet, a zero uniform floor leaves every bin + # at weight zero. Initialize that cold-start state as uniform. + sampling_weight_sum = float(sampling_probs.sum()) + if sampling_weight_sum <= 0.0: + sampling_probs = np.full(self.bin_count, 1.0 / self.bin_count, dtype=np.float64) + else: + sampling_probs = sampling_probs / sampling_weight_sum # Sample bins sampled_bins = ( diff --git a/tests/envs/test_motion_loader.py b/tests/envs/test_motion_loader.py index 1139b8ba0..ceb9c56d8 100644 --- a/tests/envs/test_motion_loader.py +++ b/tests/envs/test_motion_loader.py @@ -208,6 +208,26 @@ def test_motion_sampler_uses_env_owned_rng_and_steps_only_selected_rows(tmp_path assert sampler.current_frames[2] == frames[1] + 1 +def test_motion_sampler_adaptive_zero_floor_cold_starts_uniform(tmp_path): + motion = tmp_path / "motion.npz" + _write_motion_npz(motion, base_value=0.0, num_frames=70) + loader = MotionLoader(str(motion)) + sampler = MotionSampler( + loader, + mode="adaptive", + num_envs=64, + adaptive_uniform_ratio=0.0, + rng=np.random.default_rng(7), + ) + + frames = sampler.sample_frames(np.arange(64, dtype=np.int32)) + + assert frames.min() >= 0 + assert frames.max() < loader.num_frames + np.testing.assert_allclose(sampler.sampling_entropy, 1.0) + np.testing.assert_allclose(sampler.sampling_top1_prob, 1.0 / sampler.bin_count) + + def test_box_motion_loader_reads_object_state_and_trims_robot_joints(tmp_path): from unilab.tasks.motion_tracking.g1.motion_box_loader import BoxMotionLoader