embodichain.learning.rl.utils#
RL helper utilities: algorithm config, optimizers, and observation helpers.
Overview#
The utils package contains helper utilities for RL configuration,
data conversion, and training orchestration. It exposes the
AlgorithmCfg config and the data-conversion helpers
dict_to_tensordict() and flatten_dict_observation(), backed by the
config, helper, and trainer submodules.
Classes
Shared fields for RL algorithm configs. |
Functions
|
Convert an environment observation mapping into a TensorDict. |
Flatten a hierarchical observation TensorDict into a 2D tensor. |
Submodules
Configuration Helpers#
Classes:
Shared fields for RL algorithm configs. |
|
Optional LR scheduler. |
|
Policy optimizer configuration. |
- class embodichain.learning.rl.utils.config.AlgorithmCfg[source]#
Bases:
objectShared fields for RL algorithm configs.
Methods:
__init__([device, optimizer, lr_scheduler, ...])copy(**kwargs)Return a new object replacing specified fields with new values.
replace(**kwargs)Return a new object replacing specified fields with new values.
to_dict()Convert an object into dictionary recursively.
validate([prefix])Check the validity of configclass object.
Attributes:
- __init__(device=<factory>, optimizer=<factory>, lr_scheduler=<factory>, batch_size=<factory>, gamma=<factory>, gae_lambda=<factory>, max_grad_norm=<factory>)#
-
batch_size:
int#
- copy(**kwargs)#
Return a new object replacing specified fields with new values.
This is especially useful for frozen classes. Example usage:
@configclass(frozen=True) class C: x: int y: int c = C(1, 2) c1 = c.replace(x=3) assert c1.x == 3 and c1.y == 2
- Parameters:
obj (
object) – The object to replace.**kwargs – The fields to replace and their new values.
- Return type:
object- Returns:
The new object.
-
device:
str#
-
gae_lambda:
float#
-
gamma:
float#
-
lr_scheduler:
LRSchedulerCfg#
-
max_grad_norm:
float#
-
optimizer:
OptimizerCfg#
- replace(**kwargs)#
Return a new object replacing specified fields with new values.
This is especially useful for frozen classes. Example usage:
@configclass(frozen=True) class C: x: int y: int c = C(1, 2) c1 = c.replace(x=3) assert c1.x == 3 and c1.y == 2
- Parameters:
obj (
object) – The object to replace.**kwargs – The fields to replace and their new values.
- Return type:
object- Returns:
The new object.
- to_dict()#
Convert an object into dictionary recursively.
Note
Ignores all names starting with “__” (i.e. built-in methods).
- Parameters:
obj (
object) – An instance of a class to convert.- Raises:
ValueError – When input argument is not an object.
- Return type:
dict[str,Any]- Returns:
Converted dictionary mapping.
- validate(prefix='')#
Check the validity of configclass object.
This function checks if the object is a valid configclass object. A valid configclass object contains no MISSING entries.
- Parameters:
obj (
object) – The object to check.prefix (
str) – The prefix to add to the missing fields. Defaults to ‘’.
- Return type:
list[str]- Returns:
A list of missing fields.
- Raises:
TypeError – When the object is not a valid configuration object.
- class embodichain.learning.rl.utils.config.LRSchedulerCfg[source]#
Bases:
objectOptional LR scheduler.
name=Nonedisables scheduling.Horizon keys (
total_iters/T_max) may be omitted and bound later byBaseAlgorithm.bind_schedule.Methods:
__init__([name, kwargs])copy(**kwargs)Return a new object replacing specified fields with new values.
replace(**kwargs)Return a new object replacing specified fields with new values.
to_dict()Convert an object into dictionary recursively.
validate([prefix])Check the validity of configclass object.
Attributes:
- __init__(name=<factory>, kwargs=<factory>)#
- copy(**kwargs)#
Return a new object replacing specified fields with new values.
This is especially useful for frozen classes. Example usage:
@configclass(frozen=True) class C: x: int y: int c = C(1, 2) c1 = c.replace(x=3) assert c1.x == 3 and c1.y == 2
- Parameters:
obj (
object) – The object to replace.**kwargs – The fields to replace and their new values.
- Return type:
object- Returns:
The new object.
-
kwargs:
dict[str,Any]#
-
name:
str|None#
- replace(**kwargs)#
Return a new object replacing specified fields with new values.
This is especially useful for frozen classes. Example usage:
@configclass(frozen=True) class C: x: int y: int c = C(1, 2) c1 = c.replace(x=3) assert c1.x == 3 and c1.y == 2
- Parameters:
obj (
object) – The object to replace.**kwargs – The fields to replace and their new values.
- Return type:
object- Returns:
The new object.
- to_dict()#
Convert an object into dictionary recursively.
Note
Ignores all names starting with “__” (i.e. built-in methods).
- Parameters:
obj (
object) – An instance of a class to convert.- Raises:
ValueError – When input argument is not an object.
- Return type:
dict[str,Any]- Returns:
Converted dictionary mapping.
- validate(prefix='')#
Check the validity of configclass object.
This function checks if the object is a valid configclass object. A valid configclass object contains no MISSING entries.
- Parameters:
obj (
object) – The object to check.prefix (
str) – The prefix to add to the missing fields. Defaults to ‘’.
- Return type:
list[str]- Returns:
A list of missing fields.
- Raises:
TypeError – When the object is not a valid configuration object.
- class embodichain.learning.rl.utils.config.OptimizerCfg[source]#
Bases:
objectPolicy optimizer configuration.
Methods:
__init__([name, learning_rate, kwargs])copy(**kwargs)Return a new object replacing specified fields with new values.
replace(**kwargs)Return a new object replacing specified fields with new values.
to_dict()Convert an object into dictionary recursively.
validate([prefix])Check the validity of configclass object.
Attributes:
- __init__(name=<factory>, learning_rate=<factory>, kwargs=<factory>)#
- copy(**kwargs)#
Return a new object replacing specified fields with new values.
This is especially useful for frozen classes. Example usage:
@configclass(frozen=True) class C: x: int y: int c = C(1, 2) c1 = c.replace(x=3) assert c1.x == 3 and c1.y == 2
- Parameters:
obj (
object) – The object to replace.**kwargs – The fields to replace and their new values.
- Return type:
object- Returns:
The new object.
-
kwargs:
dict[str,Any]#
-
learning_rate:
float#
-
name:
str#
- replace(**kwargs)#
Return a new object replacing specified fields with new values.
This is especially useful for frozen classes. Example usage:
@configclass(frozen=True) class C: x: int y: int c = C(1, 2) c1 = c.replace(x=3) assert c1.x == 3 and c1.y == 2
- Parameters:
obj (
object) – The object to replace.**kwargs – The fields to replace and their new values.
- Return type:
object- Returns:
The new object.
- to_dict()#
Convert an object into dictionary recursively.
Note
Ignores all names starting with “__” (i.e. built-in methods).
- Parameters:
obj (
object) – An instance of a class to convert.- Raises:
ValueError – When input argument is not an object.
- Return type:
dict[str,Any]- Returns:
Converted dictionary mapping.
- validate(prefix='')#
Check the validity of configclass object.
This function checks if the object is a valid configclass object. A valid configclass object contains no MISSING entries.
- Parameters:
obj (
object) – The object to check.prefix (
str) – The prefix to add to the missing fields. Defaults to ‘’.
- Return type:
list[str]- Returns:
A list of missing fields.
- Raises:
TypeError – When the object is not a valid configuration object.
General Helpers#
Functions:
|
Convert an environment observation mapping into a TensorDict. |
Flatten a hierarchical observation TensorDict into a 2D tensor. |
- embodichain.learning.rl.utils.helper.dict_to_tensordict(obs_dict, device)[source]#
Convert an environment observation mapping into a TensorDict.
- Parameters:
obs_dict (
Tensor|TensorDict|Mapping[str,Any]) – Tensor or mapping returned byreset()orstep().device (
device|str) – Target device for the resulting TensorDict.
- Return type:
TensorDict- Returns:
Observation TensorDict moved onto the target device.
- embodichain.learning.rl.utils.helper.flatten_dict_observation(obs)[source]#
Flatten a hierarchical observation TensorDict into a 2D tensor.
- Parameters:
obs (
TensorDict) – Observation TensorDict with batch dimension [num_envs].- Return type:
Tensor- Returns:
Flattened observation tensor of shape [num_envs, obs_dim].
Trainer Utilities#
Classes:
Algorithm-agnostic trainer that coordinates training loop, logging, and evaluation. |
- class embodichain.learning.rl.utils.trainer.Trainer[source]#
Bases:
objectAlgorithm-agnostic trainer that coordinates training loop, logging, and evaluation.
Methods:
__init__(policy, env, algorithm, ...[, ...])save_checkpoint([path])Save policy, optimizer (when available), and trainer counters.
train(total_timesteps)- __init__(policy, env, algorithm, buffer_size, batch_size, writer, eval_freq, save_freq, checkpoint_dir, exp_name, use_wandb=True, eval_env=None, event_cfg=None, eval_event_cfg=None, num_eval_episodes=5, distributed=False, rank=0, world_size=1, eval_seed=None, best_eval_metric='eval/avg_reward', best_eval_mode='max')[source]#