Jump to content

Connect SuperML | Leeroopedia MCP: Equip your AI agents with best practices, code verification, and debugging knowledge. Powered by Leeroo — building Organizational Superintelligence. Contact us at founders@leeroo.com.

Implementation:Facebookresearch Habitat lab VER ReportWorker

From Leeroopedia
Knowledge Sources
Domains Embodied_AI, Reinforcement_Learning, Distributed_Training
Last Updated 2026-02-15 00:00 GMT

Overview

ReportWorker is a dedicated worker process in the VER training system responsible for aggregating and logging training metrics, episode statistics, timing information, and preemption decider reports across all workers.

Description

This module implements the reporting subsystem for VER (Variable Experience Rollout) training. It contains two classes:

ReportWorkerProcess is an attrs-decorated process class that handles all metric reporting. It dispatches on ReportWorkerTasks and provides the following functionality:

  • Episode statistics: Collects per-episode rewards and info metrics (from extract_scalars_from_info) as they arrive from environment workers via the episode_end task. These are accumulated per rollout and gathered across all distributed workers during logging.
  • Window episode stats: Maintains windowed running means of all episode metrics (reward, success rate, SPL, etc.) across the reward window size. Only computed on rank 0.
  • Metric logging: The log_metrics method is called after each learner update. It all-reduces the step delta and time taken across workers, gathers per-worker episode stats and learner metrics, and writes to TensorBoard:
    • Reward and per-metric scalars
    • Learner loss scalars
    • FPS (overall and windowed)
    • Preemption decider statistics (real vs expected steps/time)
    • Detailed timing breakdowns for env, policy, and learner components
  • Timing aggregation: Collects timing reports from environment workers (env_timing), inference workers (policy_timing), and the learner (learner_timing). Each category maintains its own defaultdict of WindowedRunningMean values.
  • State management: Supports state_dict and load_state_dict for checkpointing the report worker's internal state (window stats, timing stats, step counts, etc.).
  • Distributed coordination: Uses init_distrib_slurm with the gloo backend for distributed communication. Metrics are gathered/all-reduced across workers using gather_objects and all_reduce.

ReportWorker is the WorkerBase subclass providing the public API: start_collection(), state_dict(), load_state_dict(), and get_window_episode_stats(). It creates shared-memory tensors for num_steps_done and time_taken that can be read from the main process.

Usage

This module is used internally by the VER trainer to centralize all reporting. The report queue receives messages from environment workers, inference workers, and the learner. Metric logging is triggered by the learner_update task after each gradient step.

Code Reference

Source Location

Signature

class ReportWorker(WorkerBase):
    def __init__(
        self,
        mp_ctx: BaseContext,
        port: int,
        config: "DictConfig",
        report_queue: BatchedQueue,
        my_t_zero: float,
        init_num_steps=0,
        run_id=None,
    ): ...
    def start_collection(self): ...
    def state_dict(self): ...
    def load_state_dict(self, state_dict): ...
    def get_window_episode_stats(self): ...

@attr.s(auto_attribs=True)
class ReportWorkerProcess(ProcessBase):
    port: int
    config: "DictConfig"
    report_queue: BatchedQueue
    my_t_zero: float
    num_steps_done: torch.Tensor
    time_taken: torch.Tensor
    ...
    def episode_end(self, data): ...
    def num_steps_collected(self, num_steps: int): ...
    def learner_update(self, data): ...
    def start_collection(self, start_time): ...
    def log_metrics(self, writer: TensorboardWriter, learner_metrics: Dict[str, float]): ...
    def state_dict(self): ...
    def load_state_dict(self, state_dict): ...
    def run(self): ...

Import

from habitat_baselines.rl.ver.report_worker import (
    ReportWorker,
    ReportWorkerProcess,
)

I/O Contract

Inputs

Name Type Required Description
mp_ctx BaseContext Yes Multiprocessing context
port int Yes Port for distributed communication (gloo backend)
config DictConfig Yes Full Habitat baselines configuration
report_queue BatchedQueue Yes Queue receiving report messages from all worker types
my_t_zero float Yes Reference start time for this process
init_num_steps int No Initial step count when resuming training (default 0)
run_id Optional[str] No TensorBoard run ID for resuming a previous run

Outputs

Name Type Description
num_steps_done torch.Tensor Shared-memory scalar tensor tracking total training steps, readable from main process
time_taken torch.Tensor Shared-memory scalar tensor tracking total training time, readable from main process
TensorBoard logs file Reward, metrics, learner losses, FPS, timing breakdowns, and preemption decider stats written to TensorBoard
window_episode_stats Dict[str, WindowedRunningMean] Windowed running means of episode-level metrics, retrievable via get_window_episode_stats()

Usage Examples

Basic Usage

import time
import torch.multiprocessing as mp
from habitat_baselines.rl.ver.report_worker import ReportWorker
from habitat_baselines.rl.ver.queue import BatchedQueue

mp_ctx = mp.get_context("forkserver")
report_queue = BatchedQueue(mp_ctx)
my_t_zero = time.perf_counter()

report_worker = ReportWorker(
    mp_ctx,
    port=8739,
    config=config,
    report_queue=report_queue,
    my_t_zero=my_t_zero,
    init_num_steps=0,
)

# Signal start of experience collection
report_worker.start_collection()

# After learner update, log metrics (done internally by VER trainer):
# report_queue.put((ReportWorkerTasks.learner_update, learner_metrics))

# Get current episode stats
window_stats = report_worker.get_window_episode_stats()

# Checkpoint
state = report_worker.state_dict()
# ... later ...
report_worker.load_state_dict(state)

# Read shared counters from main process
print(f"Steps: {int(report_worker.num_steps_done)}")
print(f"Time: {float(report_worker.time_taken):.1f}s")

Related Pages

Page Connections

Double-click a node to navigate. Hold to expand connections.
Principle
Implementation
Heuristic
Environment