mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 10:16:09 +00:00
Fix episode pool stream benchmark pacing
This commit is contained in:
@@ -385,42 +385,74 @@ def run_pool_stream_simulation(
|
|||||||
break
|
break
|
||||||
|
|
||||||
locks = _decoder_locks(cache.manifest, resident)
|
locks = _decoder_locks(cache.manifest, resident)
|
||||||
sample_period = 1.0 / target_samples_s if target_samples_s > 0 else 0.0
|
batch_size = max(1, batch_size)
|
||||||
refill_wait_s = 0.0
|
refill_wait_s = 0.0
|
||||||
deadline_miss_s = 0.0
|
deadline_miss_s = 0.0
|
||||||
replacement_count = 0
|
replacement_count = 0
|
||||||
decoded_samples: list[tuple[int, float]] = []
|
decoded_samples: list[tuple[int, float]] = []
|
||||||
start = time.perf_counter()
|
start = time.perf_counter()
|
||||||
next_deadline = start + sample_period
|
|
||||||
|
|
||||||
for idx in range(sample_count):
|
def consume_ready_replacement() -> bool:
|
||||||
if idx > 0 and samples_per_episode > 0 and idx % samples_per_episode == 0 and pending:
|
nonlocal refill_wait_s, replacement_count
|
||||||
new_ep = pending.pop(0)
|
if not pending:
|
||||||
wait_start = time.perf_counter()
|
return False
|
||||||
cache.ensure_ready(new_ep)
|
new_ep = pending.pop(0)
|
||||||
for camera_key in cache.manifest.video_keys:
|
wait_start = time.perf_counter()
|
||||||
locks[(new_ep, camera_key)] = threading.Lock()
|
cache.ensure_ready(new_ep)
|
||||||
cache.get_decoder(new_ep, camera_key)
|
_open_resident_decoders(cache, [new_ep], decode_workers=decode_workers)
|
||||||
refill_wait_s += time.perf_counter() - wait_start
|
for camera_key in cache.manifest.video_keys:
|
||||||
old_ep = resident.pop(0)
|
locks[(new_ep, camera_key)] = threading.Lock()
|
||||||
resident_set.discard(old_ep)
|
refill_wait_s += time.perf_counter() - wait_start
|
||||||
resident.append(new_ep)
|
old_ep = resident.pop(0)
|
||||||
resident_set.add(new_ep)
|
resident_set.discard(old_ep)
|
||||||
replacement_count += 1
|
resident.append(new_ep)
|
||||||
schedule_one()
|
resident_set.add(new_ep)
|
||||||
|
replacement_count += 1
|
||||||
|
schedule_one()
|
||||||
|
return True
|
||||||
|
|
||||||
ep = rng.choice(resident)
|
def decode_batch(batch: list[tuple[int, float]], pool: ThreadPoolExecutor | None) -> None:
|
||||||
relative_t = rng.random()
|
if pool is None:
|
||||||
_decode_training_sample(cache, ep, relative_t, locks)
|
for ep, relative_t in batch:
|
||||||
decoded_samples.append((ep, relative_t))
|
_decode_training_sample(cache, ep, relative_t, locks)
|
||||||
|
return
|
||||||
|
futures = [
|
||||||
|
pool.submit(_decode_training_sample, cache, ep, relative_t, locks) for ep, relative_t in batch
|
||||||
|
]
|
||||||
|
for future in futures:
|
||||||
|
future.result()
|
||||||
|
|
||||||
if sample_period > 0:
|
samples_done = 0
|
||||||
now = time.perf_counter()
|
decode_pool = ThreadPoolExecutor(max_workers=decode_workers) if decode_workers > 1 else None
|
||||||
if now < next_deadline:
|
try:
|
||||||
time.sleep(next_deadline - now)
|
while samples_done < sample_count:
|
||||||
else:
|
batch_start = time.perf_counter()
|
||||||
deadline_miss_s += now - next_deadline
|
|
||||||
next_deadline += sample_period
|
if samples_per_episode > 0:
|
||||||
|
target_replacements = samples_done // samples_per_episode
|
||||||
|
while replacement_count < target_replacements and consume_ready_replacement():
|
||||||
|
pass
|
||||||
|
|
||||||
|
current_batch_size = min(batch_size, sample_count - samples_done)
|
||||||
|
batch = [(rng.choice(resident), rng.random()) for _ in range(current_batch_size)]
|
||||||
|
decode_batch(batch, decode_pool)
|
||||||
|
decoded_samples.extend(batch)
|
||||||
|
samples_done += current_batch_size
|
||||||
|
|
||||||
|
if samples_per_episode > 0:
|
||||||
|
target_replacements = samples_done // samples_per_episode
|
||||||
|
while replacement_count < target_replacements and consume_ready_replacement():
|
||||||
|
pass
|
||||||
|
|
||||||
|
target_batch_s = current_batch_size / target_samples_s if target_samples_s > 0 else 0.0
|
||||||
|
batch_elapsed = time.perf_counter() - batch_start
|
||||||
|
if target_batch_s > 0 and batch_elapsed < target_batch_s:
|
||||||
|
time.sleep(target_batch_s - batch_elapsed)
|
||||||
|
elif target_batch_s > 0:
|
||||||
|
deadline_miss_s += batch_elapsed - target_batch_s
|
||||||
|
finally:
|
||||||
|
if decode_pool is not None:
|
||||||
|
decode_pool.shutdown(wait=True)
|
||||||
|
|
||||||
elapsed = time.perf_counter() - start
|
elapsed = time.perf_counter() - start
|
||||||
result = {
|
result = {
|
||||||
@@ -433,6 +465,8 @@ def run_pool_stream_simulation(
|
|||||||
"replacement_episodes_s": replacement_count / elapsed if elapsed > 0 else 0.0,
|
"replacement_episodes_s": replacement_count / elapsed if elapsed > 0 else 0.0,
|
||||||
"samples_per_episode": float(samples_per_episode),
|
"samples_per_episode": float(samples_per_episode),
|
||||||
"prefetch_episodes": float(prefetch_episodes),
|
"prefetch_episodes": float(prefetch_episodes),
|
||||||
|
"batch_size": float(batch_size),
|
||||||
|
"decode_workers": float(decode_workers),
|
||||||
"kept_up": 1.0
|
"kept_up": 1.0
|
||||||
if sample_count / elapsed >= target_samples_s * 0.98 and deadline_miss_s < elapsed * 0.02
|
if sample_count / elapsed >= target_samples_s * 0.98 and deadline_miss_s < elapsed * 0.02
|
||||||
else 0.0,
|
else 0.0,
|
||||||
@@ -504,10 +538,10 @@ def _current_rss_mib() -> float | None:
|
|||||||
|
|
||||||
def _peak_rss_mib() -> float:
|
def _peak_rss_mib() -> float:
|
||||||
rss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
|
rss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
|
||||||
# Linux reports KiB; macOS reports bytes.
|
# Linux reports KiB even for very large processes; macOS reports bytes.
|
||||||
if rss > 10**8:
|
if Path("/proc/self/status").exists():
|
||||||
return rss / 1024**2
|
return rss / 1024
|
||||||
return rss / 1024
|
return rss / 1024**2
|
||||||
|
|
||||||
|
|
||||||
def _memory_snapshot() -> dict[str, float | None]:
|
def _memory_snapshot() -> dict[str, float | None]:
|
||||||
@@ -1181,6 +1215,8 @@ def run_indexed_strategy(
|
|||||||
print(f"| target samples/s | {fetch_pool['pool_stream_target_samples_s']:.1f} |")
|
print(f"| target samples/s | {fetch_pool['pool_stream_target_samples_s']:.1f} |")
|
||||||
print(f"| actual samples/s | {fetch_pool['pool_stream_actual_samples_s']:.1f} |")
|
print(f"| actual samples/s | {fetch_pool['pool_stream_actual_samples_s']:.1f} |")
|
||||||
print(f"| kept up | {'yes' if fetch_pool['pool_stream_kept_up'] else 'no'} |")
|
print(f"| kept up | {'yes' if fetch_pool['pool_stream_kept_up'] else 'no'} |")
|
||||||
|
print(f"| batch size | {fetch_pool['pool_stream_batch_size']:.0f} |")
|
||||||
|
print(f"| decode workers | {fetch_pool['pool_stream_decode_workers']:.0f} |")
|
||||||
print(f"| stream wall s | {fetch_pool['pool_stream_stream_wall_s']:.3f} |")
|
print(f"| stream wall s | {fetch_pool['pool_stream_stream_wall_s']:.3f} |")
|
||||||
print(f"| refill wait s | {fetch_pool['pool_stream_refill_wait_s']:.3f} |")
|
print(f"| refill wait s | {fetch_pool['pool_stream_refill_wait_s']:.3f} |")
|
||||||
print(f"| deadline miss s | {fetch_pool['pool_stream_deadline_miss_s']:.3f} |")
|
print(f"| deadline miss s | {fetch_pool['pool_stream_deadline_miss_s']:.3f} |")
|
||||||
|
|||||||
Reference in New Issue
Block a user