train.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348
  1. from __future__ import annotations
  2. import argparse
  3. import copy
  4. import sys
  5. from datetime import datetime
  6. from pathlib import Path
  7. from typing import Any
  8. SCRIPT_ROOT = Path(__file__).resolve().parents[1]
  9. if str(SCRIPT_ROOT) not in sys.path:
  10. sys.path.insert(0, str(SCRIPT_ROOT))
  11. import torch
  12. from guguji_rl.config import load_config, resolve_project_path, save_yaml
  13. from guguji_rl.evaluation import evaluate_forward_progress, print_forward_progress_summary
  14. def resolve_device(device_name: str) -> str:
  15. if device_name == 'auto':
  16. return 'cuda' if torch.cuda.is_available() else 'cpu'
  17. if device_name == 'cuda' and not torch.cuda.is_available():
  18. raise RuntimeError('配置要求使用 CUDA,但当前 torch 检测不到可用 GPU。')
  19. return device_name
  20. def parse_args() -> argparse.Namespace:
  21. parser = argparse.ArgumentParser(description='Train PPO policy for guguji biped robot.')
  22. parser.add_argument(
  23. '--config',
  24. default='configs/balance_ppo.yaml',
  25. help='训练配置文件路径,默认使用 balance_ppo.yaml',
  26. )
  27. parser.add_argument(
  28. '--device',
  29. default=None,
  30. help='可选覆盖配置文件中的设备设置,例如 cpu / cuda / auto',
  31. )
  32. parser.add_argument(
  33. '--total-timesteps',
  34. type=int,
  35. default=None,
  36. help='可选覆盖配置文件中的 total_timesteps',
  37. )
  38. parser.add_argument(
  39. '--init-model',
  40. default=None,
  41. help='可选指定一个已有 PPO 模型,用于继续训练或做课程学习初始化',
  42. )
  43. parser.add_argument(
  44. '--skip-auto-eval',
  45. action='store_true',
  46. help='训练完成后跳过自动前进评估',
  47. )
  48. parser.add_argument(
  49. '--render-human',
  50. action='store_true',
  51. help='如果当前后端是 MuJoCo,则在训练时同步打开 GUI 画面',
  52. )
  53. return parser.parse_args()
  54. def resolve_input_path(path_str: str) -> Path:
  55. path = Path(path_str)
  56. if path.is_absolute() or path.exists():
  57. return path
  58. return SCRIPT_ROOT / path
  59. def maybe_override_policy_log_std(model: object, initial_log_std: float | None) -> None:
  60. """可选地缩小 PPO 的初始探索方差,适合课程学习后的精修阶段。"""
  61. if initial_log_std is None:
  62. return
  63. policy = getattr(model, 'policy', None)
  64. if policy is None or not hasattr(policy, 'log_std'):
  65. raise RuntimeError('当前策略对象不支持直接设置 log_std。')
  66. # 这里直接把每个动作维度的对数标准差统一改成同一个值,
  67. # 方便在“已有步态基础上继续训练”时降低探索噪声,减少无意义的乱踢。
  68. policy.log_std.data.fill_(float(initial_log_std))
  69. print(f'已将策略初始 log_std 设为: {float(initial_log_std):.3f}')
  70. def sanitize_stage_name(stage_name: str) -> str:
  71. sanitized = ''.join(
  72. character if character.isalnum() or character in {'-', '_'} else '_'
  73. for character in stage_name.strip()
  74. )
  75. return sanitized.strip('_') or 'stage'
  76. def deep_merge_stage_override(base: dict[str, Any], override: dict[str, Any]) -> None:
  77. for key, value in override.items():
  78. if isinstance(value, dict) and isinstance(base.get(key), dict):
  79. deep_merge_stage_override(base[key], value)
  80. else:
  81. base[key] = copy.deepcopy(value)
  82. def format_stage_target_summary(stage_config: dict[str, Any]) -> str:
  83. commands_config = stage_config.get('commands', {})
  84. if bool(commands_config.get('enabled', False)):
  85. forward_velocity_range = commands_config.get('forward_velocity_range', [0.0, 0.0])
  86. yaw_rate_range = commands_config.get('yaw_rate_range', [0.0, 0.0])
  87. return (
  88. 'command_conditioned '
  89. f'vx=[{float(forward_velocity_range[0]):.2f}, {float(forward_velocity_range[1]):.2f}] '
  90. f'yaw=[{float(yaw_rate_range[0]):.2f}, {float(yaw_rate_range[1]):.2f}]'
  91. )
  92. return f'target_forward_velocity={float(stage_config["task"]["target_forward_velocity"]):.2f}'
  93. def warm_start_policy_from_checkpoint(
  94. *,
  95. ppo_class: type,
  96. model: object,
  97. init_model_path: Path,
  98. device: str,
  99. ) -> None:
  100. try:
  101. model.set_parameters(
  102. str(init_model_path),
  103. exact_match=False,
  104. device=device,
  105. )
  106. print(f'已加载课程初始化模型: {init_model_path}')
  107. return
  108. except Exception as error:
  109. print(f'完整参数加载失败,将改用兼容 warm start: {error}')
  110. source_model = ppo_class.load(str(init_model_path), device=device)
  111. source_state_dict = source_model.policy.state_dict()
  112. target_state_dict = model.policy.state_dict()
  113. matched_keys: list[str] = []
  114. skipped_keys: list[str] = []
  115. for key, source_value in source_state_dict.items():
  116. target_value = target_state_dict.get(key)
  117. if target_value is None or tuple(target_value.shape) != tuple(source_value.shape):
  118. skipped_keys.append(key)
  119. continue
  120. target_state_dict[key] = source_value.detach().clone()
  121. matched_keys.append(key)
  122. model.policy.load_state_dict(target_state_dict, strict=False)
  123. print(
  124. '已完成兼容 warm start: '
  125. f'匹配 {len(matched_keys)} 个张量,'
  126. f'跳过 {len(skipped_keys)} 个形状不兼容张量。'
  127. )
  128. if skipped_keys:
  129. print(f'跳过的典型张量: {", ".join(skipped_keys[:4])}')
  130. def build_curriculum_stage_configs(config: dict[str, Any]) -> list[tuple[str | None, dict[str, Any]]]:
  131. """把课程学习阶段展开成一组可直接训练的独立配置。"""
  132. raw_stages = config['training'].get('curriculum_stages') or []
  133. if not raw_stages:
  134. single_stage_config = copy.deepcopy(config)
  135. single_stage_config['training'].pop('curriculum_stages', None)
  136. return [(None, single_stage_config)]
  137. stage_configs: list[tuple[str | None, dict[str, Any]]] = []
  138. for stage_index, raw_stage in enumerate(raw_stages, start=1):
  139. if not isinstance(raw_stage, dict):
  140. raise RuntimeError('training.curriculum_stages 里的每个阶段都必须是字典。')
  141. stage_config = copy.deepcopy(config)
  142. stage_config['training'].pop('curriculum_stages', None)
  143. for section_name in ('task', 'commands', 'rewards', 'robot', 'sim', 'mujoco', 'evaluation', 'training'):
  144. section_override = raw_stage.get(section_name)
  145. if isinstance(section_override, dict):
  146. deep_merge_stage_override(stage_config[section_name], section_override)
  147. raw_name = str(raw_stage.get('name') or f'stage_{stage_index}')
  148. stage_name = f'{stage_index:02d}_{sanitize_stage_name(raw_name)}'
  149. # 课程阶段目前主要控制“目标前进速度 + 本阶段训练步数 + 探索方差”。
  150. # 这样 walking 阶段就能从慢到快逐段抬升,而不用一次把目标速度顶太高。
  151. if 'target_forward_velocity' in raw_stage:
  152. stage_config['task']['target_forward_velocity'] = float(raw_stage['target_forward_velocity'])
  153. if 'target_yaw_rate' in raw_stage:
  154. stage_config['task']['target_yaw_rate'] = float(raw_stage['target_yaw_rate'])
  155. if 'total_timesteps' in raw_stage:
  156. stage_config['training']['total_timesteps'] = int(raw_stage['total_timesteps'])
  157. if 'initial_log_std' in raw_stage:
  158. stage_config['training']['initial_log_std'] = float(raw_stage['initial_log_std'])
  159. stage_config['experiment']['name'] = f"{config['experiment']['name']}_{stage_name}"
  160. stage_configs.append((stage_name, stage_config))
  161. return stage_configs
  162. def main() -> int:
  163. args = parse_args()
  164. try:
  165. from stable_baselines3 import PPO
  166. from stable_baselines3.common.callbacks import CheckpointCallback
  167. from stable_baselines3.common.monitor import Monitor
  168. except ImportError:
  169. print(
  170. '缺少 stable-baselines3,请先进入 guguji_rl 目录安装依赖: '
  171. 'pip install -r requirements.txt',
  172. file=sys.stderr,
  173. )
  174. return 1
  175. from guguji_rl.envs import build_env_from_config
  176. config = load_config(resolve_input_path(args.config))
  177. if args.device is not None:
  178. config['training']['device'] = args.device
  179. if args.total_timesteps is not None:
  180. config['training']['total_timesteps'] = args.total_timesteps
  181. if args.init_model is not None:
  182. config['training']['init_model_path'] = str(resolve_input_path(args.init_model))
  183. if args.render_human:
  184. # MuJoCo 训练默认关闭渲染以保证速度。
  185. # 当你想一边训练一边看画面时,可以通过命令行临时打开 human 渲染。
  186. config = copy.deepcopy(config)
  187. config.setdefault('mujoco', {})
  188. config['mujoco']['render_mode'] = 'human'
  189. # 这里统一解析训练设备,方便你只改 YAML 就切换 CPU / GPU。
  190. config['training']['device'] = resolve_device(config['training']['device'])
  191. output_root = resolve_project_path(config, config['training']['output_root'])
  192. timestamp = datetime.now().strftime('%Y%m%d_%H%M%S')
  193. run_dir = output_root / f"{config['experiment']['name']}_{timestamp}"
  194. run_dir.mkdir(parents=True, exist_ok=True)
  195. # 保存一份展开后的配置,便于后面复现实验。
  196. save_yaml(config, run_dir / 'resolved_config.yaml')
  197. stage_configs = build_curriculum_stage_configs(config)
  198. print(f"训练设备: {config['training']['device']}")
  199. print(f"输出目录: {run_dir}")
  200. model = None
  201. final_model_path = run_dir / 'final_model'
  202. final_stage_config = config
  203. for stage_index, (stage_name, stage_config) in enumerate(stage_configs, start=1):
  204. stage_dir = run_dir if stage_name is None else run_dir / stage_name
  205. stage_dir.mkdir(parents=True, exist_ok=True)
  206. # 每个阶段都单独保存一份实际生效的配置,后面你回看实验会很方便。
  207. save_yaml(stage_config, stage_dir / 'resolved_config.yaml')
  208. if stage_name is not None:
  209. print(
  210. f'开始课程阶段 {stage_index}/{len(stage_configs)}: {stage_name} '
  211. f'({format_stage_target_summary(stage_config)}, '
  212. f'timesteps={int(stage_config["training"]["total_timesteps"])})'
  213. )
  214. env = Monitor(build_env_from_config(stage_config))
  215. checkpoint_callback = CheckpointCallback(
  216. save_freq=max(int(stage_config['training']['checkpoint_freq']), 1),
  217. save_path=str(stage_dir / 'checkpoints'),
  218. name_prefix='guguji_ppo',
  219. )
  220. try:
  221. if model is None:
  222. policy_kwargs = {
  223. 'net_arch': list(stage_config['training']['policy_net_arch']),
  224. }
  225. # 先用 MLP + PPO 跑通训练闭环,后面你可以再逐步增大网络规模。
  226. model = PPO(
  227. policy='MlpPolicy',
  228. env=env,
  229. verbose=1,
  230. seed=int(stage_config['training']['seed']),
  231. learning_rate=float(stage_config['training']['learning_rate']),
  232. n_steps=int(stage_config['training']['n_steps']),
  233. batch_size=int(stage_config['training']['batch_size']),
  234. gamma=float(stage_config['training']['gamma']),
  235. gae_lambda=float(stage_config['training']['gae_lambda']),
  236. clip_range=float(stage_config['training']['clip_range']),
  237. ent_coef=float(stage_config['training']['ent_coef']),
  238. vf_coef=float(stage_config['training']['vf_coef']),
  239. device=stage_config['training']['device'],
  240. tensorboard_log=str(run_dir / 'tensorboard'),
  241. policy_kwargs=policy_kwargs,
  242. )
  243. init_model_path = stage_config['training'].get('init_model_path')
  244. if init_model_path:
  245. resolved_init_model_path = resolve_input_path(str(init_model_path))
  246. # 如果 observation 维度还没变,这里会完整复用旧权重。
  247. # 如果我们给新任务增加了命令维度,这里会自动退化成“兼容 warm start”,
  248. # 尽量保留平衡模型里已经学到的站立/稳定控制能力。
  249. warm_start_policy_from_checkpoint(
  250. ppo_class=PPO,
  251. model=model,
  252. init_model_path=resolved_init_model_path,
  253. device=stage_config['training']['device'],
  254. )
  255. else:
  256. model.set_env(env)
  257. maybe_override_policy_log_std(model, stage_config['training'].get('initial_log_std'))
  258. model.learn(
  259. total_timesteps=int(stage_config['training']['total_timesteps']),
  260. callback=checkpoint_callback,
  261. progress_bar=True,
  262. reset_num_timesteps=(stage_index == 1),
  263. )
  264. stage_model_path = stage_dir / 'final_model'
  265. model.save(stage_model_path)
  266. final_model_path = stage_model_path
  267. final_stage_config = stage_config
  268. if stage_name is not None:
  269. print(f'课程阶段完成,模型已保存到: {stage_model_path.with_suffix(".zip")}')
  270. finally:
  271. env.close()
  272. if final_model_path != run_dir / 'final_model' and model is not None:
  273. # 在课程学习模式下,额外在 run 根目录保存一份最终模型,方便统一引用。
  274. model.save(run_dir / 'final_model')
  275. final_model_path = run_dir / 'final_model'
  276. print(f'训练完成,模型已保存到: {run_dir / "final_model.zip"}')
  277. evaluation_config = final_stage_config['evaluation']
  278. if bool(evaluation_config.get('auto_forward_progress', True)) and not args.skip_auto_eval:
  279. try:
  280. # 每轮训练结束后自动做一次前进评估,方便你快速看 delta_x / mean_vx。
  281. summary = evaluate_forward_progress(
  282. config=final_stage_config,
  283. model_path=final_model_path,
  284. episodes=int(evaluation_config['forward_progress_episodes']),
  285. max_steps=int(evaluation_config['forward_progress_max_steps']),
  286. deterministic=bool(evaluation_config['forward_progress_deterministic']),
  287. )
  288. print_forward_progress_summary(summary)
  289. except Exception as error:
  290. print(f'自动前进评估失败: {error}', file=sys.stderr)
  291. return 0
  292. if __name__ == '__main__':
  293. raise SystemExit(main())