@baimingtao/lbs-agent 1.0.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (1114) hide show
  1. package/AGENTS.md +54 -0
  2. package/README.md +181 -0
  3. package/SOUL.md +31 -0
  4. package/application/__init__.py +1 -0
  5. package/application/container.py +376 -0
  6. package/application/contracts/__init__.py +5 -0
  7. package/application/contracts/result.py +84 -0
  8. package/application/onboarding.py +243 -0
  9. package/application/runtime_factory.py +198 -0
  10. package/application/runtime_paths.py +61 -0
  11. package/application/services/__init__.py +1 -0
  12. package/application/services/agent_service.py +159 -0
  13. package/application/services/cron_service.py +49 -0
  14. package/application/services/gateway_service.py +156 -0
  15. package/application/services/interactive_service.py +275 -0
  16. package/application/services/mcp_service.py +102 -0
  17. package/application/services/model_service.py +519 -0
  18. package/application/services/orchestration_service.py +169 -0
  19. package/application/services/session_service.py +130 -0
  20. package/application/services/skill_service.py +369 -0
  21. package/application/services/tool_service.py +165 -0
  22. package/bin/lbs.js +109 -0
  23. package/config/__init__.py +2 -0
  24. package/config/config.py +453 -0
  25. package/config/config_contract.py +198 -0
  26. package/config/mcp.yaml +30 -0
  27. package/config/state_root.py +28 -0
  28. package/extensions/mcps/optional/linear/manifest.yaml +38 -0
  29. package/extensions/mcps/optional/n8n/manifest.yaml +77 -0
  30. package/extensions/mcps/optional/unreal-engine/manifest.yaml +54 -0
  31. package/extensions/prompts/agents/audit.md +46 -0
  32. package/extensions/prompts/agents/auditor.md +71 -0
  33. package/extensions/prompts/agents/browser_worker.md +56 -0
  34. package/extensions/prompts/agents/chat.md +23 -0
  35. package/extensions/prompts/agents/code.md +52 -0
  36. package/extensions/prompts/agents/data_analyst.md +39 -0
  37. package/extensions/prompts/agents/document.md +19 -0
  38. package/extensions/prompts/agents/email_agent.md +30 -0
  39. package/extensions/prompts/agents/github_agent.md +33 -0
  40. package/extensions/prompts/agents/planner.md +40 -0
  41. package/extensions/prompts/agents/research.md +50 -0
  42. package/extensions/prompts/agents/security.md +28 -0
  43. package/extensions/prompts/agents/verifier.md +28 -0
  44. package/extensions/prompts/agents/vision.md +27 -0
  45. package/extensions/skills/builtin/autonomous-ai-agents/DESCRIPTION.md +3 -0
  46. package/extensions/skills/builtin/autonomous-ai-agents/claude-code/SKILL.md +669 -0
  47. package/extensions/skills/builtin/autonomous-ai-agents/codex/SKILL.md +138 -0
  48. package/extensions/skills/builtin/autonomous-ai-agents/hermes-agent/SKILL.md +728 -0
  49. package/extensions/skills/builtin/autonomous-ai-agents/hermes-agent/references/native-mcp.md +344 -0
  50. package/extensions/skills/builtin/autonomous-ai-agents/hermes-agent/references/webhooks.md +193 -0
  51. package/extensions/skills/builtin/autonomous-ai-agents/opencode/SKILL.md +219 -0
  52. package/extensions/skills/builtin/browser-use/SKILL.md +152 -0
  53. package/extensions/skills/builtin/computer-use/SKILL.md +242 -0
  54. package/extensions/skills/builtin/creative/DESCRIPTION.md +3 -0
  55. package/extensions/skills/builtin/creative/architecture-diagram/SKILL.md +148 -0
  56. package/extensions/skills/builtin/creative/architecture-diagram/templates/template.html +319 -0
  57. package/extensions/skills/builtin/creative/ascii-art/SKILL.md +322 -0
  58. package/extensions/skills/builtin/creative/ascii-video/README.md +290 -0
  59. package/extensions/skills/builtin/creative/ascii-video/SKILL.md +241 -0
  60. package/extensions/skills/builtin/creative/ascii-video/references/architecture.md +568 -0
  61. package/extensions/skills/builtin/creative/ascii-video/references/composition.md +662 -0
  62. package/extensions/skills/builtin/creative/ascii-video/references/effects.md +640 -0
  63. package/extensions/skills/builtin/creative/ascii-video/references/inputs.md +685 -0
  64. package/extensions/skills/builtin/creative/ascii-video/references/optimization.md +688 -0
  65. package/extensions/skills/builtin/creative/ascii-video/references/scenes.md +741 -0
  66. package/extensions/skills/builtin/creative/ascii-video/references/shaders.md +683 -0
  67. package/extensions/skills/builtin/creative/ascii-video/references/troubleshooting.md +367 -0
  68. package/extensions/skills/builtin/creative/baoyu-infographic/PORT_NOTES.md +43 -0
  69. package/extensions/skills/builtin/creative/baoyu-infographic/SKILL.md +237 -0
  70. package/extensions/skills/builtin/creative/baoyu-infographic/references/analysis-framework.md +182 -0
  71. package/extensions/skills/builtin/creative/baoyu-infographic/references/base-prompt.md +43 -0
  72. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/bento-grid.md +41 -0
  73. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/binary-comparison.md +48 -0
  74. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/bridge.md +41 -0
  75. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/circular-flow.md +41 -0
  76. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/comic-strip.md +41 -0
  77. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/comparison-matrix.md +41 -0
  78. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/dashboard.md +41 -0
  79. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/dense-modules.md +72 -0
  80. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/funnel.md +41 -0
  81. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/hierarchical-layers.md +48 -0
  82. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/hub-spoke.md +41 -0
  83. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/iceberg.md +41 -0
  84. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/isometric-map.md +41 -0
  85. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/jigsaw.md +41 -0
  86. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/linear-progression.md +48 -0
  87. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/periodic-table.md +41 -0
  88. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/story-mountain.md +41 -0
  89. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/structural-breakdown.md +48 -0
  90. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/tree-branching.md +41 -0
  91. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/venn-diagram.md +41 -0
  92. package/extensions/skills/builtin/creative/baoyu-infographic/references/layouts/winding-roadmap.md +41 -0
  93. package/extensions/skills/builtin/creative/baoyu-infographic/references/structured-content-template.md +244 -0
  94. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/aged-academia.md +36 -0
  95. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/bold-graphic.md +36 -0
  96. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/chalkboard.md +61 -0
  97. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/claymation.md +29 -0
  98. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/corporate-memphis.md +29 -0
  99. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/craft-handmade.md +44 -0
  100. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/cyberpunk-neon.md +29 -0
  101. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/hand-drawn-edu.md +63 -0
  102. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/ikea-manual.md +29 -0
  103. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/kawaii.md +29 -0
  104. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/knolling.md +29 -0
  105. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/lego-brick.md +29 -0
  106. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/morandi-journal.md +60 -0
  107. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/origami.md +29 -0
  108. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/pixel-art.md +29 -0
  109. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/pop-laboratory.md +48 -0
  110. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/retro-pop-grid.md +47 -0
  111. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/storybook-watercolor.md +29 -0
  112. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/subway-map.md +29 -0
  113. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/technical-schematic.md +36 -0
  114. package/extensions/skills/builtin/creative/baoyu-infographic/references/styles/ui-wireframe.md +29 -0
  115. package/extensions/skills/builtin/creative/claude-design/SKILL.md +591 -0
  116. package/extensions/skills/builtin/creative/comfyui/SKILL.md +609 -0
  117. package/extensions/skills/builtin/creative/comfyui/references/official-cli.md +252 -0
  118. package/extensions/skills/builtin/creative/comfyui/references/rest-api.md +297 -0
  119. package/extensions/skills/builtin/creative/comfyui/references/template-integrity.md +227 -0
  120. package/extensions/skills/builtin/creative/comfyui/references/workflow-format.md +205 -0
  121. package/extensions/skills/builtin/creative/comfyui/scripts/_common.py +835 -0
  122. package/extensions/skills/builtin/creative/comfyui/scripts/auto_fix_deps.py +225 -0
  123. package/extensions/skills/builtin/creative/comfyui/scripts/check_deps.py +437 -0
  124. package/extensions/skills/builtin/creative/comfyui/scripts/comfyui_setup.sh +286 -0
  125. package/extensions/skills/builtin/creative/comfyui/scripts/extract_schema.py +315 -0
  126. package/extensions/skills/builtin/creative/comfyui/scripts/fetch_logs.py +157 -0
  127. package/extensions/skills/builtin/creative/comfyui/scripts/hardware_check.py +497 -0
  128. package/extensions/skills/builtin/creative/comfyui/scripts/health_check.py +223 -0
  129. package/extensions/skills/builtin/creative/comfyui/scripts/run_batch.py +243 -0
  130. package/extensions/skills/builtin/creative/comfyui/scripts/run_workflow.py +796 -0
  131. package/extensions/skills/builtin/creative/comfyui/scripts/ws_monitor.py +267 -0
  132. package/extensions/skills/builtin/creative/comfyui/tests/README.md +44 -0
  133. package/extensions/skills/builtin/creative/comfyui/tests/conftest.py +64 -0
  134. package/extensions/skills/builtin/creative/comfyui/tests/pytest.ini +5 -0
  135. package/extensions/skills/builtin/creative/comfyui/tests/test_check_deps.py +68 -0
  136. package/extensions/skills/builtin/creative/comfyui/tests/test_cloud_integration.py +95 -0
  137. package/extensions/skills/builtin/creative/comfyui/tests/test_common.py +443 -0
  138. package/extensions/skills/builtin/creative/comfyui/tests/test_extract_schema.py +184 -0
  139. package/extensions/skills/builtin/creative/comfyui/tests/test_run_workflow.py +210 -0
  140. package/extensions/skills/builtin/creative/comfyui/workflows/README.md +86 -0
  141. package/extensions/skills/builtin/creative/comfyui/workflows/animatediff_video.json +64 -0
  142. package/extensions/skills/builtin/creative/comfyui/workflows/flux_dev_txt2img.json +78 -0
  143. package/extensions/skills/builtin/creative/comfyui/workflows/sd15_txt2img.json +49 -0
  144. package/extensions/skills/builtin/creative/comfyui/workflows/sdxl_img2img.json +54 -0
  145. package/extensions/skills/builtin/creative/comfyui/workflows/sdxl_inpaint.json +59 -0
  146. package/extensions/skills/builtin/creative/comfyui/workflows/sdxl_txt2img.json +49 -0
  147. package/extensions/skills/builtin/creative/comfyui/workflows/upscale_4x.json +27 -0
  148. package/extensions/skills/builtin/creative/comfyui/workflows/wan_video_t2v.json +69 -0
  149. package/extensions/skills/builtin/creative/design-md/SKILL.md +171 -0
  150. package/extensions/skills/builtin/creative/design-md/templates/starter.md +92 -0
  151. package/extensions/skills/builtin/creative/excalidraw/SKILL.md +195 -0
  152. package/extensions/skills/builtin/creative/excalidraw/references/colors.md +44 -0
  153. package/extensions/skills/builtin/creative/excalidraw/references/dark-mode.md +67 -0
  154. package/extensions/skills/builtin/creative/excalidraw/references/examples.md +140 -0
  155. package/extensions/skills/builtin/creative/excalidraw/scripts/upload.py +133 -0
  156. package/extensions/skills/builtin/creative/humanizer/LICENSE +21 -0
  157. package/extensions/skills/builtin/creative/humanizer/SKILL.md +578 -0
  158. package/extensions/skills/builtin/creative/manim-video/README.md +23 -0
  159. package/extensions/skills/builtin/creative/manim-video/SKILL.md +269 -0
  160. package/extensions/skills/builtin/creative/manim-video/references/animation-design-thinking.md +161 -0
  161. package/extensions/skills/builtin/creative/manim-video/references/animations.md +282 -0
  162. package/extensions/skills/builtin/creative/manim-video/references/camera-and-3d.md +135 -0
  163. package/extensions/skills/builtin/creative/manim-video/references/decorations.md +202 -0
  164. package/extensions/skills/builtin/creative/manim-video/references/equations.md +216 -0
  165. package/extensions/skills/builtin/creative/manim-video/references/graphs-and-data.md +163 -0
  166. package/extensions/skills/builtin/creative/manim-video/references/mobjects.md +333 -0
  167. package/extensions/skills/builtin/creative/manim-video/references/paper-explainer.md +255 -0
  168. package/extensions/skills/builtin/creative/manim-video/references/production-quality.md +190 -0
  169. package/extensions/skills/builtin/creative/manim-video/references/rendering.md +185 -0
  170. package/extensions/skills/builtin/creative/manim-video/references/scene-planning.md +118 -0
  171. package/extensions/skills/builtin/creative/manim-video/references/troubleshooting.md +135 -0
  172. package/extensions/skills/builtin/creative/manim-video/references/updaters-and-trackers.md +260 -0
  173. package/extensions/skills/builtin/creative/manim-video/references/visual-design.md +124 -0
  174. package/extensions/skills/builtin/creative/manim-video/scripts/setup.sh +14 -0
  175. package/extensions/skills/builtin/creative/p5js/README.md +64 -0
  176. package/extensions/skills/builtin/creative/p5js/SKILL.md +555 -0
  177. package/extensions/skills/builtin/creative/p5js/references/animation.md +439 -0
  178. package/extensions/skills/builtin/creative/p5js/references/color-systems.md +352 -0
  179. package/extensions/skills/builtin/creative/p5js/references/core-api.md +410 -0
  180. package/extensions/skills/builtin/creative/p5js/references/export-pipeline.md +566 -0
  181. package/extensions/skills/builtin/creative/p5js/references/interaction.md +398 -0
  182. package/extensions/skills/builtin/creative/p5js/references/shapes-and-geometry.md +300 -0
  183. package/extensions/skills/builtin/creative/p5js/references/troubleshooting.md +532 -0
  184. package/extensions/skills/builtin/creative/p5js/references/typography.md +302 -0
  185. package/extensions/skills/builtin/creative/p5js/references/visual-effects.md +895 -0
  186. package/extensions/skills/builtin/creative/p5js/references/webgl-and-3d.md +423 -0
  187. package/extensions/skills/builtin/creative/p5js/scripts/export-frames.js +179 -0
  188. package/extensions/skills/builtin/creative/p5js/scripts/render.sh +108 -0
  189. package/extensions/skills/builtin/creative/p5js/scripts/serve.sh +28 -0
  190. package/extensions/skills/builtin/creative/p5js/scripts/setup.sh +87 -0
  191. package/extensions/skills/builtin/creative/p5js/templates/viewer.html +395 -0
  192. package/extensions/skills/builtin/creative/popular-web-designs/SKILL.md +202 -0
  193. package/extensions/skills/builtin/creative/popular-web-designs/templates/airbnb.md +259 -0
  194. package/extensions/skills/builtin/creative/popular-web-designs/templates/airtable.md +102 -0
  195. package/extensions/skills/builtin/creative/popular-web-designs/templates/apple.md +326 -0
  196. package/extensions/skills/builtin/creative/popular-web-designs/templates/bmw.md +193 -0
  197. package/extensions/skills/builtin/creative/popular-web-designs/templates/cal.md +272 -0
  198. package/extensions/skills/builtin/creative/popular-web-designs/templates/claude.md +325 -0
  199. package/extensions/skills/builtin/creative/popular-web-designs/templates/clay.md +317 -0
  200. package/extensions/skills/builtin/creative/popular-web-designs/templates/clickhouse.md +294 -0
  201. package/extensions/skills/builtin/creative/popular-web-designs/templates/cohere.md +279 -0
  202. package/extensions/skills/builtin/creative/popular-web-designs/templates/coinbase.md +142 -0
  203. package/extensions/skills/builtin/creative/popular-web-designs/templates/composio.md +320 -0
  204. package/extensions/skills/builtin/creative/popular-web-designs/templates/cursor.md +322 -0
  205. package/extensions/skills/builtin/creative/popular-web-designs/templates/elevenlabs.md +278 -0
  206. package/extensions/skills/builtin/creative/popular-web-designs/templates/expo.md +294 -0
  207. package/extensions/skills/builtin/creative/popular-web-designs/templates/figma.md +233 -0
  208. package/extensions/skills/builtin/creative/popular-web-designs/templates/framer.md +259 -0
  209. package/extensions/skills/builtin/creative/popular-web-designs/templates/hashicorp.md +291 -0
  210. package/extensions/skills/builtin/creative/popular-web-designs/templates/ibm.md +345 -0
  211. package/extensions/skills/builtin/creative/popular-web-designs/templates/intercom.md +159 -0
  212. package/extensions/skills/builtin/creative/popular-web-designs/templates/kraken.md +138 -0
  213. package/extensions/skills/builtin/creative/popular-web-designs/templates/linear.app.md +380 -0
  214. package/extensions/skills/builtin/creative/popular-web-designs/templates/lovable.md +311 -0
  215. package/extensions/skills/builtin/creative/popular-web-designs/templates/minimax.md +270 -0
  216. package/extensions/skills/builtin/creative/popular-web-designs/templates/mintlify.md +339 -0
  217. package/extensions/skills/builtin/creative/popular-web-designs/templates/miro.md +121 -0
  218. package/extensions/skills/builtin/creative/popular-web-designs/templates/mistral.ai.md +274 -0
  219. package/extensions/skills/builtin/creative/popular-web-designs/templates/mongodb.md +279 -0
  220. package/extensions/skills/builtin/creative/popular-web-designs/templates/notion.md +322 -0
  221. package/extensions/skills/builtin/creative/popular-web-designs/templates/nvidia.md +306 -0
  222. package/extensions/skills/builtin/creative/popular-web-designs/templates/ollama.md +280 -0
  223. package/extensions/skills/builtin/creative/popular-web-designs/templates/opencode.ai.md +294 -0
  224. package/extensions/skills/builtin/creative/popular-web-designs/templates/pinterest.md +243 -0
  225. package/extensions/skills/builtin/creative/popular-web-designs/templates/posthog.md +269 -0
  226. package/extensions/skills/builtin/creative/popular-web-designs/templates/raycast.md +281 -0
  227. package/extensions/skills/builtin/creative/popular-web-designs/templates/replicate.md +274 -0
  228. package/extensions/skills/builtin/creative/popular-web-designs/templates/resend.md +316 -0
  229. package/extensions/skills/builtin/creative/popular-web-designs/templates/revolut.md +198 -0
  230. package/extensions/skills/builtin/creative/popular-web-designs/templates/runwayml.md +257 -0
  231. package/extensions/skills/builtin/creative/popular-web-designs/templates/sanity.md +370 -0
  232. package/extensions/skills/builtin/creative/popular-web-designs/templates/sentry.md +275 -0
  233. package/extensions/skills/builtin/creative/popular-web-designs/templates/spacex.md +207 -0
  234. package/extensions/skills/builtin/creative/popular-web-designs/templates/spotify.md +259 -0
  235. package/extensions/skills/builtin/creative/popular-web-designs/templates/stripe.md +335 -0
  236. package/extensions/skills/builtin/creative/popular-web-designs/templates/supabase.md +268 -0
  237. package/extensions/skills/builtin/creative/popular-web-designs/templates/superhuman.md +265 -0
  238. package/extensions/skills/builtin/creative/popular-web-designs/templates/together.ai.md +276 -0
  239. package/extensions/skills/builtin/creative/popular-web-designs/templates/uber.md +308 -0
  240. package/extensions/skills/builtin/creative/popular-web-designs/templates/vercel.md +323 -0
  241. package/extensions/skills/builtin/creative/popular-web-designs/templates/voltagent.md +336 -0
  242. package/extensions/skills/builtin/creative/popular-web-designs/templates/warp.md +265 -0
  243. package/extensions/skills/builtin/creative/popular-web-designs/templates/webflow.md +105 -0
  244. package/extensions/skills/builtin/creative/popular-web-designs/templates/wise.md +186 -0
  245. package/extensions/skills/builtin/creative/popular-web-designs/templates/x.ai.md +270 -0
  246. package/extensions/skills/builtin/creative/popular-web-designs/templates/zapier.md +341 -0
  247. package/extensions/skills/builtin/creative/pretext/SKILL.md +220 -0
  248. package/extensions/skills/builtin/creative/pretext/references/patterns.md +258 -0
  249. package/extensions/skills/builtin/creative/pretext/templates/donut-orbit.html +1468 -0
  250. package/extensions/skills/builtin/creative/pretext/templates/hello-orb-flow.html +95 -0
  251. package/extensions/skills/builtin/creative/sketch/SKILL.md +218 -0
  252. package/extensions/skills/builtin/creative/songwriting-and-ai-music/SKILL.md +273 -0
  253. package/extensions/skills/builtin/creative/touchdesigner-mcp/SKILL.md +356 -0
  254. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/3d-scene.md +275 -0
  255. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/animation.md +221 -0
  256. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/audio-reactive.md +175 -0
  257. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/dat-scripting.md +352 -0
  258. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/external-data.md +322 -0
  259. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/geometry-comp.md +121 -0
  260. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/glsl.md +151 -0
  261. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/layout-compositor.md +131 -0
  262. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/mcp-tools.md +382 -0
  263. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/midi-osc.md +210 -0
  264. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/network-patterns.md +767 -0
  265. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/operator-tips.md +106 -0
  266. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/operators.md +239 -0
  267. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/panel-ui.md +281 -0
  268. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/particles.md +245 -0
  269. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/pitfalls.md +568 -0
  270. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/postfx.md +183 -0
  271. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/projection-mapping.md +211 -0
  272. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/python-api.md +463 -0
  273. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/replicator.md +197 -0
  274. package/extensions/skills/builtin/creative/touchdesigner-mcp/references/troubleshooting.md +244 -0
  275. package/extensions/skills/builtin/creative/touchdesigner-mcp/scripts/setup.sh +115 -0
  276. package/extensions/skills/builtin/data-science/DESCRIPTION.md +3 -0
  277. package/extensions/skills/builtin/data-science/jupyter-live-kernel/SKILL.md +152 -0
  278. package/extensions/skills/builtin/dogfood/SKILL.md +162 -0
  279. package/extensions/skills/builtin/dogfood/references/issue-taxonomy.md +109 -0
  280. package/extensions/skills/builtin/dogfood/templates/dogfood-report-template.md +86 -0
  281. package/extensions/skills/builtin/email/DESCRIPTION.md +3 -0
  282. package/extensions/skills/builtin/email/himalaya/SKILL.md +302 -0
  283. package/extensions/skills/builtin/email/himalaya/references/configuration.md +216 -0
  284. package/extensions/skills/builtin/email/himalaya/references/message-composition.md +199 -0
  285. package/extensions/skills/builtin/github/DESCRIPTION.md +3 -0
  286. package/extensions/skills/builtin/github/codebase-inspection/SKILL.md +116 -0
  287. package/extensions/skills/builtin/github/github-auth/SKILL.md +247 -0
  288. package/extensions/skills/builtin/github/github-auth/scripts/gh-env.sh +66 -0
  289. package/extensions/skills/builtin/github/github-code-review/SKILL.md +481 -0
  290. package/extensions/skills/builtin/github/github-code-review/references/review-output-template.md +74 -0
  291. package/extensions/skills/builtin/github/github-issues/SKILL.md +370 -0
  292. package/extensions/skills/builtin/github/github-issues/templates/bug-report.md +35 -0
  293. package/extensions/skills/builtin/github/github-issues/templates/feature-request.md +31 -0
  294. package/extensions/skills/builtin/github/github-pr-workflow/SKILL.md +367 -0
  295. package/extensions/skills/builtin/github/github-pr-workflow/references/ci-troubleshooting.md +183 -0
  296. package/extensions/skills/builtin/github/github-pr-workflow/references/conventional-commits.md +71 -0
  297. package/extensions/skills/builtin/github/github-pr-workflow/templates/pr-body-bugfix.md +35 -0
  298. package/extensions/skills/builtin/github/github-pr-workflow/templates/pr-body-feature.md +33 -0
  299. package/extensions/skills/builtin/github/github-repo-management/SKILL.md +516 -0
  300. package/extensions/skills/builtin/github/github-repo-management/references/github-api-cheatsheet.md +161 -0
  301. package/extensions/skills/builtin/index-cache/anthropics_skills_skills_.json +1 -0
  302. package/extensions/skills/builtin/index-cache/claude_marketplace_anthropics_skills.json +1 -0
  303. package/extensions/skills/builtin/index-cache/lobehub_index.json +1 -0
  304. package/extensions/skills/builtin/index-cache/openai_skills_skills_.json +1 -0
  305. package/extensions/skills/builtin/media/DESCRIPTION.md +3 -0
  306. package/extensions/skills/builtin/media/gif-search/SKILL.md +91 -0
  307. package/extensions/skills/builtin/media/heartmula/SKILL.md +171 -0
  308. package/extensions/skills/builtin/media/songsee/SKILL.md +83 -0
  309. package/extensions/skills/builtin/media/youtube-content/SKILL.md +75 -0
  310. package/extensions/skills/builtin/media/youtube-content/references/output-formats.md +56 -0
  311. package/extensions/skills/builtin/media/youtube-content/scripts/fetch_transcript.py +124 -0
  312. package/extensions/skills/builtin/mlops/DESCRIPTION.md +3 -0
  313. package/extensions/skills/builtin/mlops/evaluation/DESCRIPTION.md +3 -0
  314. package/extensions/skills/builtin/mlops/evaluation/lm-evaluation-harness/SKILL.md +495 -0
  315. package/extensions/skills/builtin/mlops/evaluation/lm-evaluation-harness/references/api-evaluation.md +490 -0
  316. package/extensions/skills/builtin/mlops/evaluation/lm-evaluation-harness/references/benchmark-guide.md +488 -0
  317. package/extensions/skills/builtin/mlops/evaluation/lm-evaluation-harness/references/custom-tasks.md +602 -0
  318. package/extensions/skills/builtin/mlops/evaluation/lm-evaluation-harness/references/distributed-eval.md +519 -0
  319. package/extensions/skills/builtin/mlops/evaluation/weights-and-biases/SKILL.md +592 -0
  320. package/extensions/skills/builtin/mlops/evaluation/weights-and-biases/references/artifacts.md +584 -0
  321. package/extensions/skills/builtin/mlops/evaluation/weights-and-biases/references/integrations.md +700 -0
  322. package/extensions/skills/builtin/mlops/evaluation/weights-and-biases/references/sweeps.md +846 -0
  323. package/extensions/skills/builtin/mlops/huggingface-hub/SKILL.md +81 -0
  324. package/extensions/skills/builtin/mlops/inference/DESCRIPTION.md +3 -0
  325. package/extensions/skills/builtin/mlops/inference/llama-cpp/SKILL.md +249 -0
  326. package/extensions/skills/builtin/mlops/inference/llama-cpp/references/advanced-usage.md +504 -0
  327. package/extensions/skills/builtin/mlops/inference/llama-cpp/references/hub-discovery.md +168 -0
  328. package/extensions/skills/builtin/mlops/inference/llama-cpp/references/optimization.md +89 -0
  329. package/extensions/skills/builtin/mlops/inference/llama-cpp/references/quantization.md +243 -0
  330. package/extensions/skills/builtin/mlops/inference/llama-cpp/references/server.md +150 -0
  331. package/extensions/skills/builtin/mlops/inference/llama-cpp/references/troubleshooting.md +442 -0
  332. package/extensions/skills/builtin/mlops/inference/vllm/SKILL.md +369 -0
  333. package/extensions/skills/builtin/mlops/inference/vllm/references/optimization.md +226 -0
  334. package/extensions/skills/builtin/mlops/inference/vllm/references/quantization.md +284 -0
  335. package/extensions/skills/builtin/mlops/inference/vllm/references/server-deployment.md +255 -0
  336. package/extensions/skills/builtin/mlops/inference/vllm/references/troubleshooting.md +448 -0
  337. package/extensions/skills/builtin/mlops/models/DESCRIPTION.md +3 -0
  338. package/extensions/skills/builtin/mlops/models/audiocraft/SKILL.md +568 -0
  339. package/extensions/skills/builtin/mlops/models/audiocraft/references/advanced-usage.md +666 -0
  340. package/extensions/skills/builtin/mlops/models/audiocraft/references/troubleshooting.md +504 -0
  341. package/extensions/skills/builtin/mlops/models/segment-anything/SKILL.md +506 -0
  342. package/extensions/skills/builtin/mlops/models/segment-anything/references/advanced-usage.md +589 -0
  343. package/extensions/skills/builtin/mlops/models/segment-anything/references/troubleshooting.md +484 -0
  344. package/extensions/skills/builtin/note-taking/DESCRIPTION.md +3 -0
  345. package/extensions/skills/builtin/note-taking/obsidian/SKILL.md +61 -0
  346. package/extensions/skills/builtin/productivity/DESCRIPTION.md +3 -0
  347. package/extensions/skills/builtin/productivity/airtable/SKILL.md +229 -0
  348. package/extensions/skills/builtin/productivity/google-workspace/SKILL.md +314 -0
  349. package/extensions/skills/builtin/productivity/google-workspace/references/gmail-search-syntax.md +63 -0
  350. package/extensions/skills/builtin/productivity/google-workspace/scripts/_hermes_home.py +42 -0
  351. package/extensions/skills/builtin/productivity/google-workspace/scripts/google_api.py +1225 -0
  352. package/extensions/skills/builtin/productivity/google-workspace/scripts/gws_bridge.py +111 -0
  353. package/extensions/skills/builtin/productivity/google-workspace/scripts/setup.py +481 -0
  354. package/extensions/skills/builtin/productivity/maps/SKILL.md +182 -0
  355. package/extensions/skills/builtin/productivity/maps/scripts/maps_client.py +1297 -0
  356. package/extensions/skills/builtin/productivity/nano-pdf/SKILL.md +52 -0
  357. package/extensions/skills/builtin/productivity/notion/SKILL.md +448 -0
  358. package/extensions/skills/builtin/productivity/notion/references/block-types.md +112 -0
  359. package/extensions/skills/builtin/productivity/ocr-and-documents/DESCRIPTION.md +3 -0
  360. package/extensions/skills/builtin/productivity/ocr-and-documents/SKILL.md +172 -0
  361. package/extensions/skills/builtin/productivity/ocr-and-documents/scripts/extract_marker.py +87 -0
  362. package/extensions/skills/builtin/productivity/ocr-and-documents/scripts/extract_pymupdf.py +98 -0
  363. package/extensions/skills/builtin/productivity/petdex/SKILL.md +76 -0
  364. package/extensions/skills/builtin/productivity/powerpoint/LICENSE.txt +30 -0
  365. package/extensions/skills/builtin/productivity/powerpoint/SKILL.md +237 -0
  366. package/extensions/skills/builtin/productivity/powerpoint/editing.md +205 -0
  367. package/extensions/skills/builtin/productivity/powerpoint/pptxgenjs.md +420 -0
  368. package/extensions/skills/builtin/productivity/powerpoint/scripts/__init__.py +0 -0
  369. package/extensions/skills/builtin/productivity/powerpoint/scripts/add_slide.py +195 -0
  370. package/extensions/skills/builtin/productivity/powerpoint/scripts/clean.py +286 -0
  371. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/helpers/__init__.py +0 -0
  372. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/helpers/merge_runs.py +199 -0
  373. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/helpers/simplify_redlines.py +197 -0
  374. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/pack.py +159 -0
  375. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/dml-chart.xsd +1499 -0
  376. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/dml-chartDrawing.xsd +146 -0
  377. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/dml-diagram.xsd +1085 -0
  378. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/dml-lockedCanvas.xsd +11 -0
  379. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/dml-main.xsd +3081 -0
  380. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/dml-picture.xsd +23 -0
  381. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/dml-spreadsheetDrawing.xsd +185 -0
  382. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/dml-wordprocessingDrawing.xsd +287 -0
  383. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/pml.xsd +1676 -0
  384. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/shared-additionalCharacteristics.xsd +28 -0
  385. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/shared-bibliography.xsd +144 -0
  386. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/shared-commonSimpleTypes.xsd +174 -0
  387. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/shared-customXmlDataProperties.xsd +25 -0
  388. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/shared-customXmlSchemaProperties.xsd +18 -0
  389. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/shared-documentPropertiesCustom.xsd +59 -0
  390. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/shared-documentPropertiesExtended.xsd +56 -0
  391. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/shared-documentPropertiesVariantTypes.xsd +195 -0
  392. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/shared-math.xsd +582 -0
  393. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/shared-relationshipReference.xsd +25 -0
  394. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/sml.xsd +4439 -0
  395. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/vml-main.xsd +570 -0
  396. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/vml-officeDrawing.xsd +509 -0
  397. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/vml-presentationDrawing.xsd +12 -0
  398. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/vml-spreadsheetDrawing.xsd +108 -0
  399. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/vml-wordprocessingDrawing.xsd +96 -0
  400. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/wml.xsd +3646 -0
  401. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ISO-IEC29500-4_2016/xml.xsd +116 -0
  402. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ecma/fourth-edition/opc-contentTypes.xsd +42 -0
  403. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ecma/fourth-edition/opc-coreProperties.xsd +50 -0
  404. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ecma/fourth-edition/opc-digSig.xsd +49 -0
  405. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/ecma/fourth-edition/opc-relationships.xsd +33 -0
  406. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/mce/mc.xsd +75 -0
  407. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/microsoft/wml-2010.xsd +560 -0
  408. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/microsoft/wml-2012.xsd +67 -0
  409. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/microsoft/wml-2018.xsd +14 -0
  410. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/microsoft/wml-cex-2018.xsd +20 -0
  411. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/microsoft/wml-cid-2016.xsd +13 -0
  412. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/microsoft/wml-sdtdatahash-2020.xsd +4 -0
  413. package/extensions/skills/builtin/productivity/powerpoint/scripts/office/schemas/microsoft/wml-symex-2015.xsd +8 -0
  414. package/extensions/skills/builtin/productivity/teams-meeting-pipeline/SKILL.md +116 -0
  415. package/extensions/skills/builtin/research/DESCRIPTION.md +3 -0
  416. package/extensions/skills/builtin/research/arxiv/SKILL.md +282 -0
  417. package/extensions/skills/builtin/research/arxiv/scripts/search_arxiv.py +114 -0
  418. package/extensions/skills/builtin/research/blogwatcher/SKILL.md +137 -0
  419. package/extensions/skills/builtin/research/llm-wiki/SKILL.md +449 -0
  420. package/extensions/skills/builtin/research/polymarket/SKILL.md +75 -0
  421. package/extensions/skills/builtin/research/polymarket/references/api-endpoints.md +220 -0
  422. package/extensions/skills/builtin/research/polymarket/scripts/polymarket.py +284 -0
  423. package/extensions/skills/builtin/research/research-paper-writing/SKILL.md +813 -0
  424. package/extensions/skills/builtin/research/research-paper-writing/references/autoreason-methodology.md +394 -0
  425. package/extensions/skills/builtin/research/research-paper-writing/references/checklists.md +434 -0
  426. package/extensions/skills/builtin/research/research-paper-writing/references/citation-workflow.md +564 -0
  427. package/extensions/skills/builtin/research/research-paper-writing/references/experiment-patterns.md +728 -0
  428. package/extensions/skills/builtin/research/research-paper-writing/references/human-evaluation.md +476 -0
  429. package/extensions/skills/builtin/research/research-paper-writing/references/paper-types.md +481 -0
  430. package/extensions/skills/builtin/research/research-paper-writing/references/reviewer-guidelines.md +433 -0
  431. package/extensions/skills/builtin/research/research-paper-writing/references/sources.md +191 -0
  432. package/extensions/skills/builtin/research/research-paper-writing/references/writing-guide.md +474 -0
  433. package/extensions/skills/builtin/research/research-paper-writing/templates/README.md +251 -0
  434. package/extensions/skills/builtin/research/research-paper-writing/templates/aaai2026/README.md +534 -0
  435. package/extensions/skills/builtin/research/research-paper-writing/templates/aaai2026/aaai2026-unified-supp.tex +144 -0
  436. package/extensions/skills/builtin/research/research-paper-writing/templates/aaai2026/aaai2026-unified-template.tex +952 -0
  437. package/extensions/skills/builtin/research/research-paper-writing/templates/aaai2026/aaai2026.bib +111 -0
  438. package/extensions/skills/builtin/research/research-paper-writing/templates/aaai2026/aaai2026.bst +1493 -0
  439. package/extensions/skills/builtin/research/research-paper-writing/templates/aaai2026/aaai2026.sty +315 -0
  440. package/extensions/skills/builtin/research/research-paper-writing/templates/acl/README.md +48 -0
  441. package/extensions/skills/builtin/research/research-paper-writing/templates/acl/acl.sty +312 -0
  442. package/extensions/skills/builtin/research/research-paper-writing/templates/acl/acl_latex.tex +377 -0
  443. package/extensions/skills/builtin/research/research-paper-writing/templates/acl/acl_lualatex.tex +101 -0
  444. package/extensions/skills/builtin/research/research-paper-writing/templates/acl/acl_natbib.bst +1940 -0
  445. package/extensions/skills/builtin/research/research-paper-writing/templates/acl/anthology.bib.txt +26 -0
  446. package/extensions/skills/builtin/research/research-paper-writing/templates/acl/custom.bib +70 -0
  447. package/extensions/skills/builtin/research/research-paper-writing/templates/acl/formatting.md +322 -0
  448. package/extensions/skills/builtin/research/research-paper-writing/templates/colm2025/README.md +3 -0
  449. package/extensions/skills/builtin/research/research-paper-writing/templates/colm2025/colm2025_conference.bib +11 -0
  450. package/extensions/skills/builtin/research/research-paper-writing/templates/colm2025/colm2025_conference.bst +1440 -0
  451. package/extensions/skills/builtin/research/research-paper-writing/templates/colm2025/colm2025_conference.pdf +0 -0
  452. package/extensions/skills/builtin/research/research-paper-writing/templates/colm2025/colm2025_conference.sty +218 -0
  453. package/extensions/skills/builtin/research/research-paper-writing/templates/colm2025/colm2025_conference.tex +305 -0
  454. package/extensions/skills/builtin/research/research-paper-writing/templates/colm2025/fancyhdr.sty +485 -0
  455. package/extensions/skills/builtin/research/research-paper-writing/templates/colm2025/math_commands.tex +508 -0
  456. package/extensions/skills/builtin/research/research-paper-writing/templates/colm2025/natbib.sty +1246 -0
  457. package/extensions/skills/builtin/research/research-paper-writing/templates/iclr2026/fancyhdr.sty +485 -0
  458. package/extensions/skills/builtin/research/research-paper-writing/templates/iclr2026/iclr2026_conference.bib +24 -0
  459. package/extensions/skills/builtin/research/research-paper-writing/templates/iclr2026/iclr2026_conference.bst +1440 -0
  460. package/extensions/skills/builtin/research/research-paper-writing/templates/iclr2026/iclr2026_conference.pdf +0 -0
  461. package/extensions/skills/builtin/research/research-paper-writing/templates/iclr2026/iclr2026_conference.sty +246 -0
  462. package/extensions/skills/builtin/research/research-paper-writing/templates/iclr2026/iclr2026_conference.tex +414 -0
  463. package/extensions/skills/builtin/research/research-paper-writing/templates/iclr2026/math_commands.tex +508 -0
  464. package/extensions/skills/builtin/research/research-paper-writing/templates/iclr2026/natbib.sty +1246 -0
  465. package/extensions/skills/builtin/research/research-paper-writing/templates/icml2026/algorithm.sty +79 -0
  466. package/extensions/skills/builtin/research/research-paper-writing/templates/icml2026/algorithmic.sty +201 -0
  467. package/extensions/skills/builtin/research/research-paper-writing/templates/icml2026/example_paper.bib +75 -0
  468. package/extensions/skills/builtin/research/research-paper-writing/templates/icml2026/example_paper.pdf +0 -0
  469. package/extensions/skills/builtin/research/research-paper-writing/templates/icml2026/example_paper.tex +662 -0
  470. package/extensions/skills/builtin/research/research-paper-writing/templates/icml2026/fancyhdr.sty +864 -0
  471. package/extensions/skills/builtin/research/research-paper-writing/templates/icml2026/icml2026.bst +1443 -0
  472. package/extensions/skills/builtin/research/research-paper-writing/templates/icml2026/icml2026.sty +767 -0
  473. package/extensions/skills/builtin/research/research-paper-writing/templates/icml2026/icml_numpapers.pdf +0 -0
  474. package/extensions/skills/builtin/research/research-paper-writing/templates/neurips2025/Makefile +36 -0
  475. package/extensions/skills/builtin/research/research-paper-writing/templates/neurips2025/extra_pkgs.tex +53 -0
  476. package/extensions/skills/builtin/research/research-paper-writing/templates/neurips2025/main.tex +38 -0
  477. package/extensions/skills/builtin/research/research-paper-writing/templates/neurips2025/neurips.sty +382 -0
  478. package/extensions/skills/builtin/research/web-research/SKILL.md +10 -0
  479. package/extensions/skills/builtin/smart-home/DESCRIPTION.md +3 -0
  480. package/extensions/skills/builtin/smart-home/openhue/SKILL.md +109 -0
  481. package/extensions/skills/builtin/social-media/DESCRIPTION.md +3 -0
  482. package/extensions/skills/builtin/social-media/xurl/SKILL.md +429 -0
  483. package/extensions/skills/builtin/software-development/hermes-agent-skill-authoring/SKILL.md +196 -0
  484. package/extensions/skills/builtin/software-development/node-inspect-debugger/SKILL.md +319 -0
  485. package/extensions/skills/builtin/software-development/plan/SKILL.md +338 -0
  486. package/extensions/skills/builtin/software-development/python-debugpy/SKILL.md +375 -0
  487. package/extensions/skills/builtin/software-development/requesting-code-review/SKILL.md +275 -0
  488. package/extensions/skills/builtin/software-development/simplify-code/SKILL.md +118 -0
  489. package/extensions/skills/builtin/software-development/spike/SKILL.md +197 -0
  490. package/extensions/skills/builtin/software-development/systematic-debugging/SKILL.md +411 -0
  491. package/extensions/skills/builtin/software-development/test-driven-development/SKILL.md +362 -0
  492. package/extensions/skills/builtin/yuanbao/SKILL.md +108 -0
  493. package/extensions/skills/optional/DESCRIPTION.md +22 -0
  494. package/extensions/skills/optional/autonomous-ai-agents/DESCRIPTION.md +1 -0
  495. package/extensions/skills/optional/autonomous-ai-agents/antigravity-cli/SKILL.md +197 -0
  496. package/extensions/skills/optional/autonomous-ai-agents/antigravity-cli/references/cli-docs.md +64 -0
  497. package/extensions/skills/optional/autonomous-ai-agents/blackbox/SKILL.md +144 -0
  498. package/extensions/skills/optional/autonomous-ai-agents/grok/SKILL.md +246 -0
  499. package/extensions/skills/optional/autonomous-ai-agents/honcho/SKILL.md +431 -0
  500. package/extensions/skills/optional/autonomous-ai-agents/openhands/SKILL.md +149 -0
  501. package/extensions/skills/optional/creative/baoyu-article-illustrator/PORT_NOTES.md +48 -0
  502. package/extensions/skills/optional/creative/baoyu-article-illustrator/SKILL.md +207 -0
  503. package/extensions/skills/optional/creative/baoyu-article-illustrator/prompts/system.md +32 -0
  504. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/palettes/macaron.md +33 -0
  505. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/palettes/mono-ink.md +42 -0
  506. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/palettes/neon.md +33 -0
  507. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/palettes/warm.md +32 -0
  508. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/prompt-construction.md +426 -0
  509. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/style-presets.md +80 -0
  510. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/blueprint.md +57 -0
  511. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/chalkboard.md +62 -0
  512. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/editorial.md +59 -0
  513. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/elegant.md +56 -0
  514. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/fantasy-animation.md +58 -0
  515. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/flat-doodle.md +61 -0
  516. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/flat.md +59 -0
  517. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/ink-notes.md +90 -0
  518. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/intuition-machine.md +57 -0
  519. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/minimal.md +58 -0
  520. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/nature.md +58 -0
  521. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/notion.md +58 -0
  522. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/pixel-art.md +57 -0
  523. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/playful.md +59 -0
  524. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/retro.md +59 -0
  525. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/scientific.md +59 -0
  526. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/screen-print.md +70 -0
  527. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/sketch-notes.md +56 -0
  528. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/sketch.md +57 -0
  529. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/vector-illustration.md +57 -0
  530. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/vintage.md +59 -0
  531. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/warm.md +58 -0
  532. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles/watercolor.md +58 -0
  533. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/styles.md +223 -0
  534. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/usage.md +50 -0
  535. package/extensions/skills/optional/creative/baoyu-article-illustrator/references/workflow.md +332 -0
  536. package/extensions/skills/optional/creative/baoyu-comic/PORT_NOTES.md +77 -0
  537. package/extensions/skills/optional/creative/baoyu-comic/SKILL.md +247 -0
  538. package/extensions/skills/optional/creative/baoyu-comic/references/analysis-framework.md +176 -0
  539. package/extensions/skills/optional/creative/baoyu-comic/references/art-styles/chalk.md +101 -0
  540. package/extensions/skills/optional/creative/baoyu-comic/references/art-styles/ink-brush.md +97 -0
  541. package/extensions/skills/optional/creative/baoyu-comic/references/art-styles/ligne-claire.md +75 -0
  542. package/extensions/skills/optional/creative/baoyu-comic/references/art-styles/manga.md +93 -0
  543. package/extensions/skills/optional/creative/baoyu-comic/references/art-styles/minimalist.md +84 -0
  544. package/extensions/skills/optional/creative/baoyu-comic/references/art-styles/realistic.md +89 -0
  545. package/extensions/skills/optional/creative/baoyu-comic/references/auto-selection.md +71 -0
  546. package/extensions/skills/optional/creative/baoyu-comic/references/base-prompt.md +98 -0
  547. package/extensions/skills/optional/creative/baoyu-comic/references/character-template.md +180 -0
  548. package/extensions/skills/optional/creative/baoyu-comic/references/layouts/cinematic.md +23 -0
  549. package/extensions/skills/optional/creative/baoyu-comic/references/layouts/dense.md +23 -0
  550. package/extensions/skills/optional/creative/baoyu-comic/references/layouts/four-panel.md +40 -0
  551. package/extensions/skills/optional/creative/baoyu-comic/references/layouts/mixed.md +23 -0
  552. package/extensions/skills/optional/creative/baoyu-comic/references/layouts/splash.md +23 -0
  553. package/extensions/skills/optional/creative/baoyu-comic/references/layouts/standard.md +23 -0
  554. package/extensions/skills/optional/creative/baoyu-comic/references/layouts/webtoon.md +30 -0
  555. package/extensions/skills/optional/creative/baoyu-comic/references/ohmsha-guide.md +85 -0
  556. package/extensions/skills/optional/creative/baoyu-comic/references/partial-workflows.md +106 -0
  557. package/extensions/skills/optional/creative/baoyu-comic/references/presets/concept-story.md +121 -0
  558. package/extensions/skills/optional/creative/baoyu-comic/references/presets/four-panel.md +107 -0
  559. package/extensions/skills/optional/creative/baoyu-comic/references/presets/ohmsha.md +114 -0
  560. package/extensions/skills/optional/creative/baoyu-comic/references/presets/shoujo.md +116 -0
  561. package/extensions/skills/optional/creative/baoyu-comic/references/presets/wuxia.md +110 -0
  562. package/extensions/skills/optional/creative/baoyu-comic/references/storyboard-template.md +142 -0
  563. package/extensions/skills/optional/creative/baoyu-comic/references/tones/action.md +110 -0
  564. package/extensions/skills/optional/creative/baoyu-comic/references/tones/dramatic.md +95 -0
  565. package/extensions/skills/optional/creative/baoyu-comic/references/tones/energetic.md +105 -0
  566. package/extensions/skills/optional/creative/baoyu-comic/references/tones/neutral.md +63 -0
  567. package/extensions/skills/optional/creative/baoyu-comic/references/tones/romantic.md +100 -0
  568. package/extensions/skills/optional/creative/baoyu-comic/references/tones/vintage.md +104 -0
  569. package/extensions/skills/optional/creative/baoyu-comic/references/tones/warm.md +94 -0
  570. package/extensions/skills/optional/creative/baoyu-comic/references/workflow.md +401 -0
  571. package/extensions/skills/optional/creative/blender-mcp/SKILL.md +117 -0
  572. package/extensions/skills/optional/creative/concept-diagrams/SKILL.md +362 -0
  573. package/extensions/skills/optional/creative/concept-diagrams/examples/apartment-floor-plan-conversion.md +244 -0
  574. package/extensions/skills/optional/creative/concept-diagrams/examples/automated-password-reset-flow.md +276 -0
  575. package/extensions/skills/optional/creative/concept-diagrams/examples/autonomous-llm-research-agent-flow.md +240 -0
  576. package/extensions/skills/optional/creative/concept-diagrams/examples/banana-journey-tree-to-smoothie.md +161 -0
  577. package/extensions/skills/optional/creative/concept-diagrams/examples/commercial-aircraft-structure.md +209 -0
  578. package/extensions/skills/optional/creative/concept-diagrams/examples/cpu-ooo-microarchitecture.md +236 -0
  579. package/extensions/skills/optional/creative/concept-diagrams/examples/electricity-grid-flow.md +182 -0
  580. package/extensions/skills/optional/creative/concept-diagrams/examples/feature-film-production-pipeline.md +172 -0
  581. package/extensions/skills/optional/creative/concept-diagrams/examples/hospital-emergency-department-flow.md +165 -0
  582. package/extensions/skills/optional/creative/concept-diagrams/examples/ml-benchmark-grouped-bar-chart.md +114 -0
  583. package/extensions/skills/optional/creative/concept-diagrams/examples/place-order-uml-sequence.md +325 -0
  584. package/extensions/skills/optional/creative/concept-diagrams/examples/smart-city-infrastructure.md +173 -0
  585. package/extensions/skills/optional/creative/concept-diagrams/examples/smartphone-layer-anatomy.md +154 -0
  586. package/extensions/skills/optional/creative/concept-diagrams/examples/sn2-reaction-mechanism.md +247 -0
  587. package/extensions/skills/optional/creative/concept-diagrams/examples/wind-turbine-structure.md +338 -0
  588. package/extensions/skills/optional/creative/concept-diagrams/references/dashboard-patterns.md +43 -0
  589. package/extensions/skills/optional/creative/concept-diagrams/references/infrastructure-patterns.md +144 -0
  590. package/extensions/skills/optional/creative/concept-diagrams/references/physical-shape-cookbook.md +42 -0
  591. package/extensions/skills/optional/creative/concept-diagrams/templates/template.html +174 -0
  592. package/extensions/skills/optional/creative/creative-ideation/SKILL.md +177 -0
  593. package/extensions/skills/optional/creative/creative-ideation/references/anti-slop.md +106 -0
  594. package/extensions/skills/optional/creative/creative-ideation/references/exercises.md +71 -0
  595. package/extensions/skills/optional/creative/creative-ideation/references/full-prompt-library.md +180 -0
  596. package/extensions/skills/optional/creative/creative-ideation/references/heuristics.md +85 -0
  597. package/extensions/skills/optional/creative/creative-ideation/references/method-catalog.md +88 -0
  598. package/extensions/skills/optional/creative/creative-ideation/references/methods/affinity-diagrams.md +67 -0
  599. package/extensions/skills/optional/creative/creative-ideation/references/methods/analogy-and-blending.md +83 -0
  600. package/extensions/skills/optional/creative/creative-ideation/references/methods/biomimicry.md +58 -0
  601. package/extensions/skills/optional/creative/creative-ideation/references/methods/chance-and-remix.md +75 -0
  602. package/extensions/skills/optional/creative/creative-ideation/references/methods/compression-progress.md +64 -0
  603. package/extensions/skills/optional/creative/creative-ideation/references/methods/creative-discipline.md +82 -0
  604. package/extensions/skills/optional/creative/creative-ideation/references/methods/defamiliarization.md +58 -0
  605. package/extensions/skills/optional/creative/creative-ideation/references/methods/derive-and-mapping.md +76 -0
  606. package/extensions/skills/optional/creative/creative-ideation/references/methods/first-principles.md +63 -0
  607. package/extensions/skills/optional/creative/creative-ideation/references/methods/jobs-to-be-done.md +73 -0
  608. package/extensions/skills/optional/creative/creative-ideation/references/methods/lateral-provocations.md +81 -0
  609. package/extensions/skills/optional/creative/creative-ideation/references/methods/leverage-points.md +70 -0
  610. package/extensions/skills/optional/creative/creative-ideation/references/methods/oblique-strategies.md +87 -0
  611. package/extensions/skills/optional/creative/creative-ideation/references/methods/oulipo.md +75 -0
  612. package/extensions/skills/optional/creative/creative-ideation/references/methods/pataphysics.md +64 -0
  613. package/extensions/skills/optional/creative/creative-ideation/references/methods/pattern-languages.md +78 -0
  614. package/extensions/skills/optional/creative/creative-ideation/references/methods/polya.md +77 -0
  615. package/extensions/skills/optional/creative/creative-ideation/references/methods/premortem-and-inversion.md +71 -0
  616. package/extensions/skills/optional/creative/creative-ideation/references/methods/scamper.md +63 -0
  617. package/extensions/skills/optional/creative/creative-ideation/references/methods/story-skeletons.md +100 -0
  618. package/extensions/skills/optional/creative/creative-ideation/references/methods/triz-principles.md +95 -0
  619. package/extensions/skills/optional/creative/creative-ideation/references/methods/volume-generation.md +74 -0
  620. package/extensions/skills/optional/creative/hyperframes/SKILL.md +191 -0
  621. package/extensions/skills/optional/creative/hyperframes/references/cli.md +185 -0
  622. package/extensions/skills/optional/creative/hyperframes/references/composition.md +129 -0
  623. package/extensions/skills/optional/creative/hyperframes/references/features.md +289 -0
  624. package/extensions/skills/optional/creative/hyperframes/references/gsap.md +136 -0
  625. package/extensions/skills/optional/creative/hyperframes/references/troubleshooting.md +137 -0
  626. package/extensions/skills/optional/creative/hyperframes/references/website-to-video.md +145 -0
  627. package/extensions/skills/optional/creative/hyperframes/scripts/setup.sh +135 -0
  628. package/extensions/skills/optional/creative/kanban-video-orchestrator/SKILL.md +158 -0
  629. package/extensions/skills/optional/creative/kanban-video-orchestrator/assets/brief.md.tmpl +79 -0
  630. package/extensions/skills/optional/creative/kanban-video-orchestrator/assets/setup.sh.tmpl +186 -0
  631. package/extensions/skills/optional/creative/kanban-video-orchestrator/assets/soul.md.tmpl +38 -0
  632. package/extensions/skills/optional/creative/kanban-video-orchestrator/references/examples.md +203 -0
  633. package/extensions/skills/optional/creative/kanban-video-orchestrator/references/intake.md +142 -0
  634. package/extensions/skills/optional/creative/kanban-video-orchestrator/references/kanban-setup.md +271 -0
  635. package/extensions/skills/optional/creative/kanban-video-orchestrator/references/monitoring.md +155 -0
  636. package/extensions/skills/optional/creative/kanban-video-orchestrator/references/role-archetypes.md +234 -0
  637. package/extensions/skills/optional/creative/kanban-video-orchestrator/references/tool-matrix.md +281 -0
  638. package/extensions/skills/optional/creative/kanban-video-orchestrator/scripts/bootstrap_pipeline.py +499 -0
  639. package/extensions/skills/optional/creative/kanban-video-orchestrator/scripts/monitor.py +195 -0
  640. package/extensions/skills/optional/creative/meme-generation/EXAMPLES.md +46 -0
  641. package/extensions/skills/optional/creative/meme-generation/SKILL.md +130 -0
  642. package/extensions/skills/optional/creative/meme-generation/scripts/generate_meme.py +470 -0
  643. package/extensions/skills/optional/creative/meme-generation/scripts/templates.json +97 -0
  644. package/extensions/skills/optional/creative/pixel-art/ATTRIBUTION.md +48 -0
  645. package/extensions/skills/optional/creative/pixel-art/SKILL.md +199 -0
  646. package/extensions/skills/optional/creative/pixel-art/references/palettes.md +49 -0
  647. package/extensions/skills/optional/creative/pixel-art/scripts/__init__.py +0 -0
  648. package/extensions/skills/optional/creative/pixel-art/scripts/palettes.py +167 -0
  649. package/extensions/skills/optional/creative/pixel-art/scripts/pixel_art.py +162 -0
  650. package/extensions/skills/optional/creative/pixel-art/scripts/pixel_art_video.py +345 -0
  651. package/extensions/skills/optional/devops/cli/SKILL.md +156 -0
  652. package/extensions/skills/optional/devops/cli/references/app-discovery.md +111 -0
  653. package/extensions/skills/optional/devops/cli/references/authentication.md +59 -0
  654. package/extensions/skills/optional/devops/cli/references/cli-reference.md +104 -0
  655. package/extensions/skills/optional/devops/cli/references/running-apps.md +171 -0
  656. package/extensions/skills/optional/devops/docker-management/SKILL.md +281 -0
  657. package/extensions/skills/optional/devops/hermes-s6-container-supervision/SKILL.md +178 -0
  658. package/extensions/skills/optional/devops/pinggy-tunnel/SKILL.md +309 -0
  659. package/extensions/skills/optional/devops/watchers/SKILL.md +111 -0
  660. package/extensions/skills/optional/devops/watchers/scripts/_watermark.py +148 -0
  661. package/extensions/skills/optional/devops/watchers/scripts/watch_github.py +169 -0
  662. package/extensions/skills/optional/devops/watchers/scripts/watch_http_json.py +131 -0
  663. package/extensions/skills/optional/devops/watchers/scripts/watch_rss.py +121 -0
  664. package/extensions/skills/optional/dogfood/DESCRIPTION.md +3 -0
  665. package/extensions/skills/optional/dogfood/adversarial-ux-test/SKILL.md +191 -0
  666. package/extensions/skills/optional/finance/3-statement-model/SKILL.md +430 -0
  667. package/extensions/skills/optional/finance/3-statement-model/references/formatting.md +118 -0
  668. package/extensions/skills/optional/finance/3-statement-model/references/formulas.md +292 -0
  669. package/extensions/skills/optional/finance/3-statement-model/references/sec-filings.md +125 -0
  670. package/extensions/skills/optional/finance/comps-analysis/SKILL.md +662 -0
  671. package/extensions/skills/optional/finance/dcf-model/SKILL.md +850 -0
  672. package/extensions/skills/optional/finance/dcf-model/TROUBLESHOOTING.md +40 -0
  673. package/extensions/skills/optional/finance/dcf-model/requirements.txt +7 -0
  674. package/extensions/skills/optional/finance/dcf-model/scripts/validate_dcf.py +291 -0
  675. package/extensions/skills/optional/finance/excel-author/SKILL.md +244 -0
  676. package/extensions/skills/optional/finance/excel-author/scripts/recalc.py +88 -0
  677. package/extensions/skills/optional/finance/lbo-model/SKILL.md +290 -0
  678. package/extensions/skills/optional/finance/merger-model/SKILL.md +144 -0
  679. package/extensions/skills/optional/finance/pptx-author/SKILL.md +173 -0
  680. package/extensions/skills/optional/finance/stocks/SKILL.md +88 -0
  681. package/extensions/skills/optional/finance/stocks/scripts/stocks_client.py +755 -0
  682. package/extensions/skills/optional/mcp/DESCRIPTION.md +3 -0
  683. package/extensions/skills/optional/mcp/fastmcp/SKILL.md +300 -0
  684. package/extensions/skills/optional/mcp/fastmcp/references/fastmcp-cli.md +110 -0
  685. package/extensions/skills/optional/mcp/fastmcp/scripts/scaffold_fastmcp.py +56 -0
  686. package/extensions/skills/optional/mcp/fastmcp/templates/api_wrapper.py +54 -0
  687. package/extensions/skills/optional/mcp/fastmcp/templates/database_server.py +77 -0
  688. package/extensions/skills/optional/mcp/fastmcp/templates/file_processor.py +55 -0
  689. package/extensions/skills/optional/mcp/mcporter/SKILL.md +123 -0
  690. package/extensions/skills/optional/migration/DESCRIPTION.md +1 -0
  691. package/extensions/skills/optional/migration/openclaw-migration/SKILL.md +298 -0
  692. package/extensions/skills/optional/migration/openclaw-migration/scripts/openclaw_to_hermes.py +3136 -0
  693. package/extensions/skills/optional/mlops/accelerate/SKILL.md +333 -0
  694. package/extensions/skills/optional/mlops/accelerate/references/custom-plugins.md +453 -0
  695. package/extensions/skills/optional/mlops/accelerate/references/megatron-integration.md +489 -0
  696. package/extensions/skills/optional/mlops/accelerate/references/performance.md +525 -0
  697. package/extensions/skills/optional/mlops/chroma/SKILL.md +408 -0
  698. package/extensions/skills/optional/mlops/chroma/references/integration.md +38 -0
  699. package/extensions/skills/optional/mlops/clip/SKILL.md +255 -0
  700. package/extensions/skills/optional/mlops/clip/references/applications.md +207 -0
  701. package/extensions/skills/optional/mlops/faiss/SKILL.md +223 -0
  702. package/extensions/skills/optional/mlops/faiss/references/index_types.md +280 -0
  703. package/extensions/skills/optional/mlops/flash-attention/SKILL.md +364 -0
  704. package/extensions/skills/optional/mlops/flash-attention/references/benchmarks.md +215 -0
  705. package/extensions/skills/optional/mlops/flash-attention/references/transformers-integration.md +293 -0
  706. package/extensions/skills/optional/mlops/guidance/SKILL.md +574 -0
  707. package/extensions/skills/optional/mlops/guidance/references/backends.md +553 -0
  708. package/extensions/skills/optional/mlops/guidance/references/constraints.md +674 -0
  709. package/extensions/skills/optional/mlops/guidance/references/examples.md +767 -0
  710. package/extensions/skills/optional/mlops/huggingface-tokenizers/SKILL.md +518 -0
  711. package/extensions/skills/optional/mlops/huggingface-tokenizers/references/algorithms.md +653 -0
  712. package/extensions/skills/optional/mlops/huggingface-tokenizers/references/integration.md +637 -0
  713. package/extensions/skills/optional/mlops/huggingface-tokenizers/references/pipeline.md +723 -0
  714. package/extensions/skills/optional/mlops/huggingface-tokenizers/references/training.md +565 -0
  715. package/extensions/skills/optional/mlops/inference/outlines/SKILL.md +654 -0
  716. package/extensions/skills/optional/mlops/inference/outlines/references/backends.md +615 -0
  717. package/extensions/skills/optional/mlops/inference/outlines/references/examples.md +773 -0
  718. package/extensions/skills/optional/mlops/inference/outlines/references/json_generation.md +649 -0
  719. package/extensions/skills/optional/mlops/instructor/SKILL.md +742 -0
  720. package/extensions/skills/optional/mlops/instructor/references/examples.md +107 -0
  721. package/extensions/skills/optional/mlops/instructor/references/providers.md +70 -0
  722. package/extensions/skills/optional/mlops/instructor/references/validation.md +606 -0
  723. package/extensions/skills/optional/mlops/lambda-labs/SKILL.md +549 -0
  724. package/extensions/skills/optional/mlops/lambda-labs/references/advanced-usage.md +611 -0
  725. package/extensions/skills/optional/mlops/lambda-labs/references/troubleshooting.md +530 -0
  726. package/extensions/skills/optional/mlops/llava/SKILL.md +306 -0
  727. package/extensions/skills/optional/mlops/llava/references/training.md +197 -0
  728. package/extensions/skills/optional/mlops/modal/SKILL.md +345 -0
  729. package/extensions/skills/optional/mlops/modal/references/advanced-usage.md +503 -0
  730. package/extensions/skills/optional/mlops/modal/references/troubleshooting.md +494 -0
  731. package/extensions/skills/optional/mlops/nemo-curator/SKILL.md +384 -0
  732. package/extensions/skills/optional/mlops/nemo-curator/references/deduplication.md +87 -0
  733. package/extensions/skills/optional/mlops/nemo-curator/references/filtering.md +102 -0
  734. package/extensions/skills/optional/mlops/obliteratus/SKILL.md +342 -0
  735. package/extensions/skills/optional/mlops/obliteratus/references/analysis-modules.md +160 -0
  736. package/extensions/skills/optional/mlops/obliteratus/references/methods-guide.md +139 -0
  737. package/extensions/skills/optional/mlops/obliteratus/templates/abliteration-config.yaml +33 -0
  738. package/extensions/skills/optional/mlops/obliteratus/templates/analysis-study.yaml +40 -0
  739. package/extensions/skills/optional/mlops/obliteratus/templates/batch-abliteration.yaml +41 -0
  740. package/extensions/skills/optional/mlops/peft/SKILL.md +435 -0
  741. package/extensions/skills/optional/mlops/peft/references/advanced-usage.md +514 -0
  742. package/extensions/skills/optional/mlops/peft/references/troubleshooting.md +480 -0
  743. package/extensions/skills/optional/mlops/pinecone/SKILL.md +360 -0
  744. package/extensions/skills/optional/mlops/pinecone/references/deployment.md +181 -0
  745. package/extensions/skills/optional/mlops/pytorch-fsdp/SKILL.md +85 -0
  746. package/extensions/skills/optional/mlops/pytorch-fsdp/references/index.md +7 -0
  747. package/extensions/skills/optional/mlops/pytorch-fsdp/references/other.md +3297 -0
  748. package/extensions/skills/optional/mlops/pytorch-lightning/SKILL.md +348 -0
  749. package/extensions/skills/optional/mlops/pytorch-lightning/references/callbacks.md +436 -0
  750. package/extensions/skills/optional/mlops/pytorch-lightning/references/distributed.md +490 -0
  751. package/extensions/skills/optional/mlops/pytorch-lightning/references/hyperparameter-tuning.md +556 -0
  752. package/extensions/skills/optional/mlops/qdrant/SKILL.md +497 -0
  753. package/extensions/skills/optional/mlops/qdrant/references/advanced-usage.md +648 -0
  754. package/extensions/skills/optional/mlops/qdrant/references/troubleshooting.md +631 -0
  755. package/extensions/skills/optional/mlops/research/DESCRIPTION.md +3 -0
  756. package/extensions/skills/optional/mlops/research/dspy/SKILL.md +592 -0
  757. package/extensions/skills/optional/mlops/research/dspy/references/examples.md +663 -0
  758. package/extensions/skills/optional/mlops/research/dspy/references/modules.md +475 -0
  759. package/extensions/skills/optional/mlops/research/dspy/references/optimizers.md +566 -0
  760. package/extensions/skills/optional/mlops/saelens/SKILL.md +390 -0
  761. package/extensions/skills/optional/mlops/saelens/references/README.md +69 -0
  762. package/extensions/skills/optional/mlops/saelens/references/api.md +333 -0
  763. package/extensions/skills/optional/mlops/saelens/references/tutorials.md +318 -0
  764. package/extensions/skills/optional/mlops/simpo/SKILL.md +220 -0
  765. package/extensions/skills/optional/mlops/simpo/references/datasets.md +478 -0
  766. package/extensions/skills/optional/mlops/simpo/references/hyperparameters.md +452 -0
  767. package/extensions/skills/optional/mlops/simpo/references/loss-functions.md +350 -0
  768. package/extensions/skills/optional/mlops/slime/SKILL.md +467 -0
  769. package/extensions/skills/optional/mlops/slime/references/api-reference.md +392 -0
  770. package/extensions/skills/optional/mlops/slime/references/troubleshooting.md +386 -0
  771. package/extensions/skills/optional/mlops/stable-diffusion/SKILL.md +523 -0
  772. package/extensions/skills/optional/mlops/stable-diffusion/references/advanced-usage.md +716 -0
  773. package/extensions/skills/optional/mlops/stable-diffusion/references/troubleshooting.md +555 -0
  774. package/extensions/skills/optional/mlops/tensorrt-llm/SKILL.md +189 -0
  775. package/extensions/skills/optional/mlops/tensorrt-llm/references/multi-gpu.md +298 -0
  776. package/extensions/skills/optional/mlops/tensorrt-llm/references/optimization.md +242 -0
  777. package/extensions/skills/optional/mlops/tensorrt-llm/references/serving.md +470 -0
  778. package/extensions/skills/optional/mlops/torchtitan/SKILL.md +361 -0
  779. package/extensions/skills/optional/mlops/torchtitan/references/checkpoint.md +181 -0
  780. package/extensions/skills/optional/mlops/torchtitan/references/custom-models.md +258 -0
  781. package/extensions/skills/optional/mlops/torchtitan/references/float8.md +133 -0
  782. package/extensions/skills/optional/mlops/torchtitan/references/fsdp.md +126 -0
  783. package/extensions/skills/optional/mlops/training/axolotl/SKILL.md +164 -0
  784. package/extensions/skills/optional/mlops/training/axolotl/references/api.md +1535 -0
  785. package/extensions/skills/optional/mlops/training/axolotl/references/dataset-formats.md +843 -0
  786. package/extensions/skills/optional/mlops/training/axolotl/references/index.md +15 -0
  787. package/extensions/skills/optional/mlops/training/axolotl/references/other.md +2367 -0
  788. package/extensions/skills/optional/mlops/training/trl-fine-tuning/SKILL.md +460 -0
  789. package/extensions/skills/optional/mlops/training/trl-fine-tuning/references/dpo-variants.md +227 -0
  790. package/extensions/skills/optional/mlops/training/trl-fine-tuning/references/grpo-training.md +504 -0
  791. package/extensions/skills/optional/mlops/training/trl-fine-tuning/references/online-rl.md +82 -0
  792. package/extensions/skills/optional/mlops/training/trl-fine-tuning/references/reward-modeling.md +122 -0
  793. package/extensions/skills/optional/mlops/training/trl-fine-tuning/references/sft-training.md +168 -0
  794. package/extensions/skills/optional/mlops/training/trl-fine-tuning/templates/basic_grpo_training.py +228 -0
  795. package/extensions/skills/optional/mlops/training/unsloth/SKILL.md +81 -0
  796. package/extensions/skills/optional/mlops/training/unsloth/references/index.md +7 -0
  797. package/extensions/skills/optional/mlops/training/unsloth/references/llms-full.md +381 -0
  798. package/extensions/skills/optional/mlops/training/unsloth/references/llms-txt.md +382 -0
  799. package/extensions/skills/optional/mlops/training/unsloth/references/llms.md +82 -0
  800. package/extensions/skills/optional/mlops/whisper/SKILL.md +319 -0
  801. package/extensions/skills/optional/mlops/whisper/references/languages.md +189 -0
  802. package/extensions/skills/optional/productivity/canvas/SKILL.md +98 -0
  803. package/extensions/skills/optional/productivity/canvas/scripts/canvas_api.py +160 -0
  804. package/extensions/skills/optional/productivity/here-now/SKILL.md +217 -0
  805. package/extensions/skills/optional/productivity/here-now/scripts/drive.sh +406 -0
  806. package/extensions/skills/optional/productivity/here-now/scripts/publish.sh +445 -0
  807. package/extensions/skills/optional/productivity/memento-flashcards/SKILL.md +324 -0
  808. package/extensions/skills/optional/productivity/memento-flashcards/scripts/memento_cards.py +353 -0
  809. package/extensions/skills/optional/productivity/memento-flashcards/scripts/youtube_quiz.py +88 -0
  810. package/extensions/skills/optional/productivity/shop/SKILL.md +224 -0
  811. package/extensions/skills/optional/productivity/shop/references/catalog-mcp.md +236 -0
  812. package/extensions/skills/optional/productivity/shop/references/direct-api.md +265 -0
  813. package/extensions/skills/optional/productivity/shop/references/legal.md +3 -0
  814. package/extensions/skills/optional/productivity/shop/references/safety.md +36 -0
  815. package/extensions/skills/optional/productivity/shopify/SKILL.md +373 -0
  816. package/extensions/skills/optional/productivity/siyuan/SKILL.md +298 -0
  817. package/extensions/skills/optional/productivity/telephony/SKILL.md +418 -0
  818. package/extensions/skills/optional/productivity/telephony/scripts/telephony.py +1343 -0
  819. package/extensions/skills/optional/research/bioinformatics/SKILL.md +235 -0
  820. package/extensions/skills/optional/research/darwinian-evolver/SKILL.md +161 -0
  821. package/extensions/skills/optional/research/darwinian-evolver/scripts/parrot_openrouter.py +218 -0
  822. package/extensions/skills/optional/research/darwinian-evolver/scripts/show_snapshot.py +92 -0
  823. package/extensions/skills/optional/research/darwinian-evolver/templates/custom_problem_template.py +240 -0
  824. package/extensions/skills/optional/research/domain-intel/SKILL.md +97 -0
  825. package/extensions/skills/optional/research/domain-intel/scripts/domain_intel.py +397 -0
  826. package/extensions/skills/optional/research/drug-discovery/SKILL.md +218 -0
  827. package/extensions/skills/optional/research/drug-discovery/references/ADMET_REFERENCE.md +66 -0
  828. package/extensions/skills/optional/research/drug-discovery/scripts/chembl_target.py +53 -0
  829. package/extensions/skills/optional/research/drug-discovery/scripts/ro5_screen.py +44 -0
  830. package/extensions/skills/optional/research/duckduckgo-search/SKILL.md +238 -0
  831. package/extensions/skills/optional/research/duckduckgo-search/scripts/duckduckgo.sh +28 -0
  832. package/extensions/skills/optional/research/gitnexus-explorer/SKILL.md +196 -0
  833. package/extensions/skills/optional/research/gitnexus-explorer/scripts/proxy.mjs +92 -0
  834. package/extensions/skills/optional/research/osint-investigation/SKILL.md +227 -0
  835. package/extensions/skills/optional/research/osint-investigation/references/sources/courtlistener.md +92 -0
  836. package/extensions/skills/optional/research/osint-investigation/references/sources/gdelt.md +94 -0
  837. package/extensions/skills/optional/research/osint-investigation/references/sources/icij-offshore.md +79 -0
  838. package/extensions/skills/optional/research/osint-investigation/references/sources/nyc-acris.md +85 -0
  839. package/extensions/skills/optional/research/osint-investigation/references/sources/ofac-sdn.md +84 -0
  840. package/extensions/skills/optional/research/osint-investigation/references/sources/opencorporates.md +88 -0
  841. package/extensions/skills/optional/research/osint-investigation/references/sources/sec-edgar.md +80 -0
  842. package/extensions/skills/optional/research/osint-investigation/references/sources/senate-ld.md +85 -0
  843. package/extensions/skills/optional/research/osint-investigation/references/sources/usaspending.md +90 -0
  844. package/extensions/skills/optional/research/osint-investigation/references/sources/wayback.md +87 -0
  845. package/extensions/skills/optional/research/osint-investigation/references/sources/wikipedia.md +96 -0
  846. package/extensions/skills/optional/research/osint-investigation/scripts/_http.py +82 -0
  847. package/extensions/skills/optional/research/osint-investigation/scripts/_normalize.py +67 -0
  848. package/extensions/skills/optional/research/osint-investigation/scripts/build_findings.py +221 -0
  849. package/extensions/skills/optional/research/osint-investigation/scripts/entity_resolution.py +228 -0
  850. package/extensions/skills/optional/research/osint-investigation/scripts/fetch_courtlistener.py +149 -0
  851. package/extensions/skills/optional/research/osint-investigation/scripts/fetch_gdelt.py +161 -0
  852. package/extensions/skills/optional/research/osint-investigation/scripts/fetch_icij_offshore.py +234 -0
  853. package/extensions/skills/optional/research/osint-investigation/scripts/fetch_nyc_acris.py +203 -0
  854. package/extensions/skills/optional/research/osint-investigation/scripts/fetch_ofac_sdn.py +175 -0
  855. package/extensions/skills/optional/research/osint-investigation/scripts/fetch_opencorporates.py +191 -0
  856. package/extensions/skills/optional/research/osint-investigation/scripts/fetch_sec_edgar.py +184 -0
  857. package/extensions/skills/optional/research/osint-investigation/scripts/fetch_senate_ld.py +146 -0
  858. package/extensions/skills/optional/research/osint-investigation/scripts/fetch_usaspending.py +170 -0
  859. package/extensions/skills/optional/research/osint-investigation/scripts/fetch_wayback.py +142 -0
  860. package/extensions/skills/optional/research/osint-investigation/scripts/fetch_wikipedia.py +266 -0
  861. package/extensions/skills/optional/research/osint-investigation/scripts/timing_analysis.py +252 -0
  862. package/extensions/skills/optional/research/osint-investigation/templates/source-template.md +57 -0
  863. package/extensions/skills/optional/research/parallel-cli/SKILL.md +391 -0
  864. package/extensions/skills/optional/research/qmd/SKILL.md +417 -0
  865. package/extensions/skills/optional/research/scrapling/SKILL.md +336 -0
  866. package/extensions/skills/optional/research/searxng-search/SKILL.md +212 -0
  867. package/extensions/skills/optional/research/searxng-search/scripts/searxng.sh +22 -0
  868. package/extensions/skills/optional/security/DESCRIPTION.md +3 -0
  869. package/extensions/skills/optional/software-development/code-wiki/SKILL.md +442 -0
  870. package/extensions/skills/optional/software-development/code-wiki/templates/README.md +31 -0
  871. package/extensions/skills/optional/software-development/code-wiki/templates/architecture.md +30 -0
  872. package/extensions/skills/optional/software-development/code-wiki/templates/getting-started.md +47 -0
  873. package/extensions/skills/optional/software-development/code-wiki/templates/module.md +38 -0
  874. package/extensions/skills/optional/software-development/rest-graphql-debug/SKILL.md +514 -0
  875. package/extensions/skills/optional/software-development/subagent-driven-development/SKILL.md +352 -0
  876. package/extensions/skills/optional/software-development/subagent-driven-development/references/context-budget-discipline.md +53 -0
  877. package/extensions/skills/optional/software-development/subagent-driven-development/references/gates-taxonomy.md +93 -0
  878. package/extensions/skills/optional/web-development/DESCRIPTION.md +5 -0
  879. package/extensions/skills/optional/web-development/cloudflare-temporary-deploy/SKILL.md +127 -0
  880. package/extensions/skills/optional/web-development/cloudflare-temporary-deploy/scripts/parse_deploy_output.py +122 -0
  881. package/extensions/skills/optional/web-development/page-agent/SKILL.md +190 -0
  882. package/interfaces/__init__.py +1 -0
  883. package/interfaces/api/__init__.py +1 -0
  884. package/interfaces/api/_impl/_models.py +105 -0
  885. package/interfaces/api/_impl/server.py +1480 -0
  886. package/interfaces/cli/__init__.py +1 -0
  887. package/interfaces/cli/_impl/__init__.py +1 -0
  888. package/interfaces/cli/_impl/commands/__init__.py +1 -0
  889. package/interfaces/cli/_impl/commands/agents.py +106 -0
  890. package/interfaces/cli/_impl/commands/architecture.py +63 -0
  891. package/interfaces/cli/_impl/commands/chat_repl.py +291 -0
  892. package/interfaces/cli/_impl/commands/common.py +17 -0
  893. package/interfaces/cli/_impl/commands/cron.py +145 -0
  894. package/interfaces/cli/_impl/commands/doctor.py +201 -0
  895. package/interfaces/cli/_impl/commands/gateway.py +111 -0
  896. package/interfaces/cli/_impl/commands/llm.py +113 -0
  897. package/interfaces/cli/_impl/commands/models.py +578 -0
  898. package/interfaces/cli/_impl/commands/orchestrator.py +71 -0
  899. package/interfaces/cli/_impl/commands/run_agent.py +102 -0
  900. package/interfaces/cli/_impl/commands/sessions.py +237 -0
  901. package/interfaces/cli/_impl/commands/skill.py +121 -0
  902. package/interfaces/cli/_impl/commands/smoke_test.py +193 -0
  903. package/interfaces/cli/_impl/commands/soul.py +106 -0
  904. package/interfaces/cli/_impl/commands/tools.py +249 -0
  905. package/interfaces/cli/_impl/handlers/__init__.py +1 -0
  906. package/interfaces/cli/_impl/main.py +727 -0
  907. package/interfaces/cli/_impl/parser/__init__.py +1 -0
  908. package/interfaces/cli/_impl/renderers/__init__.py +23 -0
  909. package/interfaces/cli/_impl/result.py +81 -0
  910. package/interfaces/gateway/__init__.py +1 -0
  911. package/interfaces/gateway/_impl/__init__.py +163 -0
  912. package/interfaces/gateway/_impl/adapters/dingtalk/__init__.py +8 -0
  913. package/interfaces/gateway/_impl/adapters/dingtalk/adapter.py +100 -0
  914. package/interfaces/gateway/_impl/adapters/feishu/__init__.py +8 -0
  915. package/interfaces/gateway/_impl/adapters/feishu/adapter.py +575 -0
  916. package/interfaces/gateway/_impl/adapters/qq/__init__.py +8 -0
  917. package/interfaces/gateway/_impl/adapters/qq/adapter.py +266 -0
  918. package/interfaces/gateway/_impl/adapters/slack/__init__.py +8 -0
  919. package/interfaces/gateway/_impl/adapters/slack/adapter.py +72 -0
  920. package/interfaces/gateway/_impl/adapters/wechat/__init__.py +8 -0
  921. package/interfaces/gateway/_impl/adapters/wechat/adapter.py +239 -0
  922. package/interfaces/gateway/_impl/middleware/__init__.py +1 -0
  923. package/interfaces/gateway/_impl/routes/__init__.py +1 -0
  924. package/interfaces/web/__init__.py +1 -0
  925. package/interfaces/web/frontend/.oxlintrc.json +8 -0
  926. package/interfaces/web/frontend/README.md +32 -0
  927. package/interfaces/web/frontend/index.html +13 -0
  928. package/interfaces/web/frontend/package-lock.json +5179 -0
  929. package/interfaces/web/frontend/package.json +48 -0
  930. package/interfaces/web/frontend/postcss.config.js +6 -0
  931. package/interfaces/web/frontend/public/cat-spritesheet.webp +0 -0
  932. package/interfaces/web/frontend/public/favicon.svg +1 -0
  933. package/interfaces/web/frontend/public/icons.svg +24 -0
  934. package/interfaces/web/frontend/run-5173.cmd +3 -0
  935. package/interfaces/web/frontend/src/App.css +184 -0
  936. package/interfaces/web/frontend/src/App.tsx +272 -0
  937. package/interfaces/web/frontend/src/api/client.ts +115 -0
  938. package/interfaces/web/frontend/src/assets/hero.png +0 -0
  939. package/interfaces/web/frontend/src/assets/react.svg +1 -0
  940. package/interfaces/web/frontend/src/assets/vite.svg +1 -0
  941. package/interfaces/web/frontend/src/components/AgentAvatar.tsx +59 -0
  942. package/interfaces/web/frontend/src/components/CatSpirit.tsx +248 -0
  943. package/interfaces/web/frontend/src/components/ChatFlow.tsx +587 -0
  944. package/interfaces/web/frontend/src/components/DateTimePicker.tsx +170 -0
  945. package/interfaces/web/frontend/src/components/MarkdownView.tsx +20 -0
  946. package/interfaces/web/frontend/src/components/NiceSelect.tsx +166 -0
  947. package/interfaces/web/frontend/src/components/PagePanel.tsx +68 -0
  948. package/interfaces/web/frontend/src/components/PixelCat.tsx +150 -0
  949. package/interfaces/web/frontend/src/components/RichEditor.tsx +129 -0
  950. package/interfaces/web/frontend/src/components/SakuraParticles.tsx +70 -0
  951. package/interfaces/web/frontend/src/components/SidebarLeft.tsx +100 -0
  952. package/interfaces/web/frontend/src/components/SidebarRight.tsx +333 -0
  953. package/interfaces/web/frontend/src/components/Toggle.tsx +19 -0
  954. package/interfaces/web/frontend/src/contexts/ThemeContext.tsx +230 -0
  955. package/interfaces/web/frontend/src/contexts/ThemeContextCore.ts +14 -0
  956. package/interfaces/web/frontend/src/contexts/useTheme.ts +8 -0
  957. package/interfaces/web/frontend/src/hooks/useAgent.ts +507 -0
  958. package/interfaces/web/frontend/src/hooks/useDashboardData.ts +160 -0
  959. package/interfaces/web/frontend/src/index.css +1903 -0
  960. package/interfaces/web/frontend/src/main.tsx +11 -0
  961. package/interfaces/web/frontend/src/mock/data.ts +3 -0
  962. package/interfaces/web/frontend/src/pages/AgentsPage.tsx +156 -0
  963. package/interfaces/web/frontend/src/pages/CronPage.tsx +80 -0
  964. package/interfaces/web/frontend/src/pages/FilesPage.tsx +57 -0
  965. package/interfaces/web/frontend/src/pages/McpPage.tsx +143 -0
  966. package/interfaces/web/frontend/src/pages/ModelsPage.tsx +170 -0
  967. package/interfaces/web/frontend/src/pages/SessionsPage.tsx +135 -0
  968. package/interfaces/web/frontend/src/pages/SkillsPage.tsx +4 -0
  969. package/interfaces/web/frontend/src/pages/SoulsPage.tsx +109 -0
  970. package/interfaces/web/frontend/src/pages/ToolsPage.tsx +3 -0
  971. package/interfaces/web/frontend/src/pages/UsagePage.tsx +132 -0
  972. package/interfaces/web/frontend/src/pages/WorkersPage.tsx +198 -0
  973. package/interfaces/web/frontend/src/types/index.ts +138 -0
  974. package/interfaces/web/frontend/tailwind.config.js +52 -0
  975. package/interfaces/web/frontend/tsconfig.app.json +26 -0
  976. package/interfaces/web/frontend/tsconfig.json +7 -0
  977. package/interfaces/web/frontend/tsconfig.node.json +23 -0
  978. package/interfaces/web/frontend/vite.config.ts +16 -0
  979. package/interfaces/web/frontend/yarn.lock +2281 -0
  980. package/interfaces/web/legacy_web/index.html +1259 -0
  981. package/lbs.py +6 -0
  982. package/package.json +53 -0
  983. package/requirements.txt +98 -0
  984. package/runtime/__init__.py +1 -0
  985. package/runtime/context/__init__.py +26 -0
  986. package/runtime/context/compression/strategies/__init__.py +1 -0
  987. package/runtime/context/compression/token_counter.py +62 -0
  988. package/runtime/context/compressor.py +594 -0
  989. package/runtime/context/management/__init__.py +1 -0
  990. package/runtime/context/management/context_builder.py +67 -0
  991. package/runtime/context/management/context_manager.py +1008 -0
  992. package/runtime/context/trimming/__init__.py +1 -0
  993. package/runtime/context/trimming/priority.py +58 -0
  994. package/runtime/context/trimming/window.py +56 -0
  995. package/runtime/core/__init__.py +1 -0
  996. package/runtime/core/llm/__init__.py +24 -0
  997. package/runtime/core/llm/auxiliary_client.py +279 -0
  998. package/runtime/core/llm/builtins.py +300 -0
  999. package/runtime/core/llm/client/__init__.py +699 -0
  1000. package/runtime/core/llm/credential.py +334 -0
  1001. package/runtime/core/llm/error_classifier.py +158 -0
  1002. package/runtime/core/llm/factory/__init__.py +1 -0
  1003. package/runtime/core/llm/health.py +106 -0
  1004. package/runtime/core/llm/model_router.py +592 -0
  1005. package/runtime/core/llm/provider.py +120 -0
  1006. package/runtime/core/llm/redact.py +142 -0
  1007. package/runtime/core/llm/registry/__init__.py +50 -0
  1008. package/runtime/core/llm/registry/models.yaml +233 -0
  1009. package/runtime/core/llm/think_scrubber.py +232 -0
  1010. package/runtime/core/llm/tokenizer.py +150 -0
  1011. package/runtime/core/llm/usage_pricing.py +373 -0
  1012. package/runtime/core/sub_agents/__init__.py +11 -0
  1013. package/runtime/core/sub_agents/agents.yaml +477 -0
  1014. package/runtime/core/sub_agents/base.py +28 -0
  1015. package/runtime/core/sub_agents/builtin/__init__.py +1 -0
  1016. package/runtime/core/sub_agents/registry.py +96 -0
  1017. package/runtime/core/tokenizer.py +146 -0
  1018. package/runtime/gateway.py +14 -0
  1019. package/runtime/memory/__init__.py +11 -0
  1020. package/runtime/memory/context_memory.py +249 -0
  1021. package/runtime/memory/long_term/knowledge_graph/__init__.py +1 -0
  1022. package/runtime/memory/long_term/sqlite/__init__.py +104 -0
  1023. package/runtime/memory/long_term/vector_store/__init__.py +1 -0
  1024. package/runtime/memory/manager.py +144 -0
  1025. package/runtime/memory/provider.py +97 -0
  1026. package/runtime/memory/providers/long_term.py +143 -0
  1027. package/runtime/memory/providers/working.py +53 -0
  1028. package/runtime/memory/session/__init__.py +4 -0
  1029. package/runtime/memory/session/identity.py +48 -0
  1030. package/runtime/memory/session/session_context.py +23 -0
  1031. package/runtime/memory/session/session_store.py +1118 -0
  1032. package/runtime/memory/short_term/__init__.py +8 -0
  1033. package/runtime/memory/short_term/buffer.py +8 -0
  1034. package/runtime/memory/short_term/working.py +93 -0
  1035. package/runtime/memory/task_contracts/__init__.py +115 -0
  1036. package/runtime/orchestration/__init__.py +8 -0
  1037. package/runtime/orchestration/auto_cycle/__init__.py +486 -0
  1038. package/runtime/orchestration/auto_cycle/trigger.py +8 -0
  1039. package/runtime/orchestration/background_review.py +470 -0
  1040. package/runtime/orchestration/display.py +206 -0
  1041. package/runtime/orchestration/iteration_budget.py +98 -0
  1042. package/runtime/orchestration/loop/__init__.py +17 -0
  1043. package/runtime/orchestration/loop/completion_gate.py +156 -0
  1044. package/runtime/orchestration/loop/context_window.py +71 -0
  1045. package/runtime/orchestration/loop/main_loop.py +2014 -0
  1046. package/runtime/orchestration/loop/message_history.py +71 -0
  1047. package/runtime/orchestration/loop/progress_tracker.py +211 -0
  1048. package/runtime/orchestration/loop/run_trace.py +70 -0
  1049. package/runtime/orchestration/loop/tool_call_executor.py +214 -0
  1050. package/runtime/orchestration/loop/tool_execution_contract.py +134 -0
  1051. package/runtime/orchestration/loop/types.py +282 -0
  1052. package/runtime/orchestration/retry_engine.py +241 -0
  1053. package/runtime/orchestration/runtime/__init__.py +20 -0
  1054. package/runtime/orchestration/runtime/arbiter.py +85 -0
  1055. package/runtime/orchestration/runtime/checkpoint.py +148 -0
  1056. package/runtime/orchestration/runtime/error_policy.py +55 -0
  1057. package/runtime/orchestration/runtime/events.py +92 -0
  1058. package/runtime/orchestration/runtime/process_registry.py +605 -0
  1059. package/runtime/orchestration/runtime/resume_manager.py +283 -0
  1060. package/runtime/orchestration/runtime/retry_policy.py +26 -0
  1061. package/runtime/orchestration/runtime/status.py +60 -0
  1062. package/runtime/orchestration/runtime/types.py +74 -0
  1063. package/runtime/orchestration/scheduler/__init__.py +8 -0
  1064. package/runtime/orchestration/scheduler/master_scheduler.py +36 -0
  1065. package/runtime/orchestration/scheduler/task_graph.py +67 -0
  1066. package/runtime/orchestration/supervisor_loop.py +1041 -0
  1067. package/runtime/orchestration/tool_guardrails.py +249 -0
  1068. package/runtime/orchestration/tracer.py +403 -0
  1069. package/runtime/orchestration/verification/__init__.py +14 -0
  1070. package/runtime/orchestration/verification/completion.py +248 -0
  1071. package/runtime/orchestration/verification/task_acceptance.py +87 -0
  1072. package/runtime/orchestration/verification/tool_result.py +38 -0
  1073. package/runtime/orchestration/verification/worker_result.py +54 -0
  1074. package/runtime/perception/README.md +24 -0
  1075. package/runtime/perception/__init__.py +5 -0
  1076. package/runtime/perception/detectors/__init__.py +1 -0
  1077. package/runtime/perception/detectors/entity.py +86 -0
  1078. package/runtime/perception/detectors/intent.py +103 -0
  1079. package/runtime/perception/input/multimodal/__init__.py +1 -0
  1080. package/runtime/perception/input/streaming/__init__.py +1 -0
  1081. package/runtime/perception/input/text/__init__.py +1 -0
  1082. package/runtime/perception/preprocess/__init__.py +1 -0
  1083. package/scripts/check_llm_layer.py +16 -0
  1084. package/scripts/check_orchestration_layer.py +15 -0
  1085. package/scripts/clean-pycache.js +56 -0
  1086. package/scripts/init.js +187 -0
  1087. package/tools/__init__.py +1 -0
  1088. package/tools/mcp/clients/__init__.py +368 -0
  1089. package/tools/mcp/registry.py +538 -0
  1090. package/tools/mcp/servers/__init__.py +46 -0
  1091. package/tools/skill/__init__.py +4 -0
  1092. package/tools/skill/curator.py +368 -0
  1093. package/tools/skill/registry/__init__.py +11 -0
  1094. package/tools/skill/registry/registry.py +367 -0
  1095. package/tools/skill/registry/source_config.py +101 -0
  1096. package/tools/skill/skills/__init__.py +1 -0
  1097. package/tools/tool/__init__.py +12 -0
  1098. package/tools/tool/advanced_tools.py +307 -0
  1099. package/tools/tool/approval_workflow.py +216 -0
  1100. package/tools/tool/builtin.py +656 -0
  1101. package/tools/tool/config.py +191 -0
  1102. package/tools/tool/hermes_migrated.py +190 -0
  1103. package/tools/tool/implementations/__init__.py +1 -0
  1104. package/tools/tool/implementations/browser_use_provider.py +691 -0
  1105. package/tools/tool/implementations/extended_tools.py +260 -0
  1106. package/tools/tool/mcp_provider.py +95 -0
  1107. package/tools/tool/policy.py +65 -0
  1108. package/tools/tool/providers/base.py +115 -0
  1109. package/tools/tool/registry/__init__.py +183 -0
  1110. package/tools/tool/runtime.py +87 -0
  1111. package/tools/tool/schema.py +99 -0
  1112. package/tools/tool/skill_provider.py +164 -0
  1113. package/tools/tool/todo_provider.py +161 -0
  1114. package/tools/tool/tools.yaml +115 -0
@@ -0,0 +1,3297 @@
1
+ # Pytorch-Fsdp - 其他
2
+
3
+ **页数:** 15
4
+
5
+ ---
6
+
7
+ ## 分布式数据并行#
8
+
9
+ **URL:** https://pytorch.org/docs/stable/notes/ddp.html
10
+
11
+ **目录:**
12
+ - 分布式数据并行#
13
+ - 示例#
14
+ - 内部设计#
15
+ - 实现#
16
+ - ProcessGroup#
17
+ - DistributedDataParallel#
18
+ - TorchDynamo DDPOptimizer#
19
+
20
+ 创建时间:2020年1月15日 | 最后更新时间:2024年1月25日
21
+
22
+ torch.nn.parallel.DistributedDataParallel 的实现随着时间不断演进。本设计说明基于 v1.4 版本的状态编写。
23
+
24
+ torch.nn.parallel.DistributedDataParallel (DDP) 透明地执行分布式数据并行训练。本页介绍了它的工作原理并揭示了实现细节。
25
+
26
+ 让我们从一个简单的 torch.nn.parallel.DistributedDataParallel 示例开始。此示例使用 torch.nn.Linear 作为本地模型,用 DDP 包装它,然后在 DDP 模型上运行一次前向传递、一次反向传递和一个优化器步骤。之后,本地模型上的参数将被更新,并且不同进程上的所有模型应该完全相同。
27
+
28
+ DDP 可以与 TorchDynamo 配合使用。与 TorchDynamo 一起使用时,请在编译模型之前应用 DDP 模型包装器,以便 torchdynamo 可以基于 DDP 存储桶(bucket)大小应用 DDPOptimizer(计算图断点优化)。(有关更多信息,请参见 TorchDynamo DDPOptimizer。)
29
+
30
+ 本节通过深入探讨一次迭代中每个步骤的细节,揭示了 torch.nn.parallel.DistributedDataParallel 背后的工作原理。
31
+
32
+ 前提条件:DDP 依赖 c10d ProcessGroup 进行通信。因此,应用程序必须在构建 DDP 之前创建 ProcessGroup 实例。
33
+
34
+ 构建:DDP 构造函数接收对本地模块的引用,并将 state_dict() 从 rank 为 0 的进程广播到组中的所有其他进程,以确保所有模型副本从完全相同的状态开始。然后,每个 DDP 进程创建一个本地 Reducer,它稍后将在反向传递期间负责梯度同步。为了提高通信效率,Reducer 将参数梯度组织到存储桶中,并每次减少(reduce)一个存储桶。可以通过在 DDP 构造函数中设置 bucket_cap_mb 参数来配置存储桶大小。参数梯度到存储桶的映射是在构建时根据存储桶大小限制和参数大小确定的。模型参数以给定模型的 Model.parameters() 的(大致)逆序分配到存储桶中。使用逆序的原因是 DDP 预期在反向传递期间梯度大约按照该顺序准备就绪。下图显示了一个示例。请注意,grad0 和 grad1 在 bucket1 中,而另外两个梯度在 bucket0 中。当然,这个假设可能并不总是成立的,当这种情况发生时,可能会降低 DDP 反向传递的速度,因为 Reducer 无法在最早可能的时间启动通信。除了分桶之外,Reducer 还会在构建期间注册自动求导钩子,每个参数一个钩子。这些钩子将在反向传递期间梯度准备就绪时被触发。
35
+
36
+ 前向传递:DDP 接收输入并将其传递给本地模型,然后如果 find_unused_parameters 设置为 True,则会分析本地模型的输出。此模式允许在模型的子图上运行反向传递,DDP 通过从模型输出遍历自动求导图并将所有未使用的参数标记为已准备好进行归约(reduction)来找出哪些参数参与了反向传递。在反向传递期间,Reducer 只会等待未准备好的参数,但它仍会归约所有存储桶。将参数梯度标记为就绪目前并不能帮助 DDP 跳过存储桶,但它可以防止 DDP 在反向传递期间永远等待缺失的梯度。请注意,遍历自动求导图会引入额外的开销,因此应用程序应仅在必要时将 find_unused_parameters 设置为 True。
37
+
38
+ 反向传递:backward() 函数直接在损失 Tensor 上调用,这超出了 DDP 的控制范围,DDP 使用在构建时注册的自动求导钩子来触发梯度同步。当一个梯度准备就绪时,在该梯度累加器上相应的 DDP 钩子将触发,然后 DDP 会将该参数梯度标记为已准备好进行归约。当一个存储桶中的所有梯度都准备就绪时,Reducer 会启动该存储桶上的异步 allreduce 操作,以计算所有进程间梯度的平均值。当所有存储桶都准备就绪时,Reducer 将阻塞等待所有 allreduce 操作完成。完成后,平均梯度将被写入所有参数的 param.grad 字段。因此,在反向传递之后,不同 DDP 进程中相同对应参数上的 grad 字段应该是相同的。
39
+
40
+ 优化器步骤:从优化器的角度来看,它正在优化一个本地模型。所有 DDP 进程上的模型副本可以保持同步,因为它们都从相同的状态开始,并且在每次迭代中都具有相同的平均梯度。
41
+
42
+ DDP 要求所有进程上的 Reducer 实例以完全相同的顺序调用 allreduce,这是通过始终按照存储桶索引顺序而不是实际的存储桶就绪顺序运行 allreduce 来实现的。跨进程的 allreduce 顺序不匹配可能导致错误结果或 DDP 反向传递挂起。
43
+
44
+ 以下是指向 DDP 实现组件的指针。堆叠图显示了代码的结构。
45
+
46
+ ProcessGroup.hpp:包含所有进程组实现的抽象 API。c10d 库提供了 3 种开箱即用的实现,即 ProcessGroupGloo、ProcessGroupNCCL 和 ProcessGroupMPI。DistributedDataParallel 在初始化期间使用 ProcessGroup::broadcast() 将模型状态从 rank 为 0 的进程发送给其他进程,并使用 ProcessGroup::allreduce() 对梯度进行求和。
47
+
48
+ Store.hpp:协助进程组实例互相发现的会合服务。
49
+
50
+ distributed.py:是 DDP 的 Python 入口点。它实现了初始化步骤和 nn.parallel.DistributedDataParallel 模块的 forward 函数,这些函数调用了 C++ 库。当一个 DDP 进程在多个设备上工作时,它的 _sync_param 函数会执行进程内参数同步,并且它还会将模型缓冲区从 rank 为 0 的进程广播到所有其他进程。进程间的参数同步发生在 Reducer.cpp 中。
51
+
52
+ comm.h:实现了合并广播辅助函数,该函数在初始化期间被调用来广播模型状态,并在前向传递之前同步模型缓冲区。
53
+
54
+ reducer.h:提供了反向传递中梯度同步的核心实现。它有三个入口点函数:
55
+
56
+ Reducer:在 distributed.py 中调用的构造函数,它将 Reducer::autograd_hook() 注册到梯度累加器。
57
+
58
+ autograd_hook() 函数将在梯度准备就绪时由自动求导引擎调用。
59
+
60
+ prepare_for_backward() 在 distributed.py 中的 DDP 前向传递结束时被调用。如果在 DDP 构造函数中 find_unused_parameters 设置为 True,它会遍历自动求导图以查找未使用的参数。
61
+
62
+ DDP 的性能优势源于在反向传递期间将 allreduce 集合通信与计算重叠起来。当与 TorchDynamo 一起使用以编译整个前向和整个反向计算图时,AotAutograd 阻止了这种重叠,因为 allreduce 操作是在整个优化后的反向计算完成_之后_由自动求导钩子启动的。
63
+
64
+ TorchDynamo 的 DDPOptimizer 通过在反向传递期间 DDP 的 allreduce 存储桶的逻辑边界处断开前向计算图来提供帮助。注意:目标是在反向传递期间断开计算图,最简单的实现是断开前向计算图,然后对每个部分调用 AotAutograd 和编译。这允许 DDP 的 allreduce 钩子在反向传递的各个部分之间触发,并调度通信以与计算重叠。
65
+
66
+ 有关更深入的解释和实验结果,请参阅此博客文章,或阅读 torch/_dynamo/optimizations/distributed.py 中的文档和代码
67
+
68
+ 要调试 DDPOptimizer,请设置 TORCH_LOGS='ddp_graphs' 以获取完整的计算图转储。对于不包含计算图的日志,请将 'dynamo'、'distributed' 或 'dist_ddp' 中的任何一个添加到 TORCH_LOGS 中(以获取有关存储桶边界的基本信息)。要禁用 DDPOptimizer,请设置 torch._dynamo.config.optimize_ddp=False。在没有 DDPOptimizer 的情况下,DDP 和 TorchDynamo 仍然应该能够正常工作,但性能会有所下降。
69
+
70
+ ---
71
+
72
+ ## PyTorch 文档#
73
+
74
+ **URL:** https://pytorch.org/docs/stable/
75
+
76
+ **目录:**
77
+ - PyTorch 文档#
78
+ - 索引和表格#
79
+
80
+ PyTorch 是一个针对使用 GPU 和 CPU 的深度学习进行了优化的张量库。
81
+
82
+ 本档中描述的特性按发布状态进行分类:
83
+
84
+ 稳定版 (API-Stable):这些功能将得到长期维护,并且通常在性能上没有重大限制或在文档上没有缺失。我们也期望保持向后兼容性(尽管可能会发生破坏性更改,并会提前一个版本发出通知)。
85
+
86
+ 不稳定版 (API-Unstable):包含所有正在积极开发中的功能,这些 API 可能会根据用户反馈、必要的性能改进或由于跨算子的覆盖范围尚未完成而发生变化。这些功能的 API 和性能特征可能会发生变化。
87
+
88
+ ---
89
+
90
+ ## 通用 Join 上下文管理器#
91
+
92
+ **URL:** https://pytorch.org/docs/stable/distributed.algorithms.join.html
93
+
94
+ **目录:**
95
+ - 通用 Join 上下文管理器#
96
+
97
+ 创建时间:2025年6月6日 | 最后更新时间:2025年6月6日
98
+
99
+ 通用 Join 上下文管理器有助于针对不均匀输入进行分布式训练。本页概述了相关类的 API:Join、Joinable 和 JoinHook。有关教程,请参阅使用 Join 上下文管理器进行不均匀输入的分布式训练。
100
+
101
+ 此类定义了通用的 join 上下文管理器,允许在某个进程加入(join)后调用自定义钩子。
102
+
103
+ 这些钩子应该掩盖(shadow)未加入进程的集合通信,以防止挂起和报错,并确保算法的正确性。有关钩子定义的详细信息,请参阅 JoinHook。
104
+
105
+ 上下文管理器要求每个参与的 Joinable 在其自身的每次迭代的集合通信之前调用 notify_join_context() 方法,以确保正确性。
106
+
107
+ 上下文管理器要求 JoinHook 对象中的所有 process_group 属性必须相同。如果有多个 JoinHook 对象,则使用第一个对象的设备。进程组和设备信息用于检查是否存在未加入的进程,并在启用 throw_on_early_termination 时通知进程抛出异常,这两者都使用 all-reduce 操作。
108
+
109
+ joinables (List[Joinable]) – 参与的 Joinable 列表;它们的钩子将按给定顺序进行迭代。
110
+
111
+ enable (bool) – 启用不均匀输入检测的标志;设置为 False 会禁用上下文管理器的功能,只能在用户确信输入不会不均匀时设置(默认:True)。
112
+
113
+ throw_on_early_termination (bool) – 控制在检测到不均匀输入时是否抛出异常的标志(默认:False)。
114
+
115
+ 通知 join 上下文管理器,调用进程尚未加入。
116
+
117
+ 然后,如果 throw_on_early_termination=True,则检查是否检测到了不均匀输入(即是否有进程已经加入),如果是,则抛出异常。
118
+
119
+ 此方法应在其每次迭代的集合通信之前从 Joinable 对象中调用。例如,在 DistributedDataParallel 中,应在前向传递开始时调用此方法。
120
+
121
+ 只有传入上下文管理器的第一个 Joinable 对象会在此方法中执行集合通信,而对于其他对象,此方法是空操作。
122
+
123
+ joinable (Joinable) – 调用此方法的 Joinable 对象。
124
+
125
+ 如果 joinable 是传入上下文管理器的第一个对象,则返回一个用于通知上下文管理器该进程尚未加入的 all-reduce 异步工作句柄;否则返回 None。
126
+
127
+ 这定义了可加入类的抽象基类。
128
+
129
+ 一个可加入类(继承自 Joinable)应该实现返回 JoinHook 实例的 join_hook(),此外还要分别返回设备和进程组信息的 join_device() 和 join_process_group()。
130
+
131
+ 返回执行 join 上下文管理器所需集合通信的设备。
132
+
133
+ 返回给定 Joinable 的 JoinHook 实例。
134
+
135
+ kwargs (dict) – 一个包含在运行时修改 join 钩子行为的任何关键字参数的字典;所有共享同一个 join 上下文管理器的 Joinable 实例都会接收到相同的 kwargs 值。
136
+
137
+ 返回 join 上下文管理器本身所需集合通信的进程组。
138
+
139
+ 这定义了一个 join 钩子,它在 join 上下文管理器中提供了两个入口点。
140
+
141
+ 入口点:一个主钩子,只要存在未加入的进程就会被重复调用;一个后置钩子,在所有进程都加入后调用一次。
142
+
143
+ 要为通用 join 上下文管理器实现 join 钩子,请定义一个继承自 JoinHook 的类,并根据需要覆盖 main_hook() 和 post_hook()。
144
+
145
+ 只要存在未加入的进程,就调用此钩子以掩盖一次训练迭代中的集合通信。
146
+
147
+ 训练迭代,即在一次前向传递、反向传递和优化器步骤中。
148
+
149
+ 在所有进程都加入后调用此钩子。
150
+
151
+ 它被传递了一个额外的布尔参数 is_last_joiner,指示该 rank 是否是最后加入的之一。
152
+
153
+ is_last_joiner (bool) – 如果该 rank 是最后加入的之一,则为 True;否则为 False。
154
+
155
+ ---
156
+
157
+ ## 实验性面向对象的分布式 API#
158
+
159
+ **URL:** https://pytorch.org/docs/stable/distributed._dist2.html
160
+
161
+ **目录:**
162
+ - 实验性面向对象的分布式 API#
163
+
164
+ 创建时间:2025年7月9日 | 最后更新时间:2025年7月30日
165
+
166
+ 这是一个用于 PyTorch 分布式的全新实验性 API。目前正在积极开发中,可能会发生变更或被完全删除。
167
+
168
+ 这旨在作为更加灵活和面向对象的分布式 API 的试验场。
169
+
170
+ 基类:pybind11_object
171
+
172
+ ProcessGroup 是一种通信原语,允许在一组进程之间进行集合操作。
173
+
174
+ 这是一个为所有 ProcessGroup 提供接口的基类。它不适合直接使用,而应由子类扩展。
175
+
176
+ 基类:pybind11_object
177
+
178
+ 用于进程组的后端类型。
179
+
180
+ 如果后端支持,中止所有操作和连接
181
+
182
+ allgather(self: torch._C._distributed_c10d.ProcessGroup, output_tensors: collections.abc.Sequence[collections.abc.Sequence[torch.Tensor]], input_tensors: collections.abc.Sequence[torch.Tensor], opts: torch._C._distributed_c10d.AllgatherOptions = <torch._C._distributed_c10d.AllgatherOptions object at 0x7f0162b6b9b0>) -> c10d::Work
183
+
184
+ 跨进程组从所有进程中 allgather(全收集)输入张量。
185
+
186
+ 有关详细信息,请参见 torch.distributed.all_gather()。
187
+
188
+ allgather(self: torch._C._distributed_c10d.ProcessGroup, output_tensors: collections.abc.Sequence[torch.Tensor], input_tensor: torch.Tensor, timeout: datetime.timedelta | None = None) -> c10d::Work
189
+
190
+ 跨进程组从所有进程中 allgather(全收集)输入张量。
191
+
192
+ 有关详细信息,请参见 torch.distributed.all_gather()。
193
+
194
+ 跨进程组从所有进程中 allgather(全收集)输入张量。
195
+
196
+ 有关详细信息,请参见 torch.distributed.all_gather()。
197
+
198
+ 跨进程组从所有进程中 allgather(全收集)输入张量。
199
+
200
+ 有关详细信息,请参见 torch.distributed.all_gather()。
201
+
202
+ allreduce(self: torch._C._distributed_c10d.ProcessGroup, tensors: collections.abc.Sequence[torch.Tensor], opts: torch._C._distributed_c10d.AllreduceOptions = <torch._C._distributed_c10d.AllreduceOptions object at 0x7f0162745db0>) -> c10d::Work
203
+
204
+ 跨进程组中的所有进程对提供的张量进行 allreduce(全归约)。
205
+
206
+ 有关详细信息,请参见 torch.distributed.all_reduce()。
207
+
208
+ allreduce(self: torch._C._distributed_c10d.ProcessGroup, tensors: collections.abc.Sequence[torch.Tensor], op: torch._C._distributed_c10d.ReduceOp = <RedOpType.SUM: 0>, timeout: datetime.timedelta | None = None) -> c10d::Work
209
+
210
+ 跨进程组中的所有进程对提供的张量进行 allreduce(全归约)。
211
+
212
+ 有关详细信息,请参见 torch.distributed.all_reduce()。
213
+
214
+ allreduce(self: torch._C._distributed_c10d.ProcessGroup, tensor: torch.Tensor, op: torch._C._distributed_c10d.ReduceOp = <RedOpType.SUM: 0>, timeout: datetime.timedelta | None = None) -> c10d::Work
215
+
216
+ 跨进程组中的所有进程对提供的张量进行 allreduce(全归约)。
217
+
218
+ 有关详细信息,请参见 torch.distributed.all_reduce()。
219
+
220
+ 跨进程组中的所有进程对提供的张量进行 allreduce(全归约)。
221
+
222
+ 有关详细信息,请参见 torch.distributed.all_reduce()。
223
+
224
+ 跨进程组从所有进程中进行 input tensor(输入张量)的 alltoall(全交换)。
225
+
226
+ 有关详细信息,请参见 torch.distributed.all_to_all()。
227
+
228
+ alltoall_base(self: torch._C._distributed_c10d.ProcessGroup, output: torch.Tensor, input: torch.Tensor, output_split_sizes: collections.abc.Sequence[typing.SupportsInt], input_split_sizes: collections.abc.Sequence[typing.SupportsInt], opts: torch._C._distributed_c10d.AllToAllOptions = <torch._C._distributed_c10d.AllToAllOptions object at 0x7f0162b79d30>) -> c10d::Work
229
+
230
+ 跨进程组从所有进程中 alltoall(全交换)输入张量。
231
+
232
+ 有关详细信息,请参见 torch.distributed.all_to_all()。
233
+
234
+ alltoall_base(self: torch._C._distributed_c10d.ProcessGroup, output: torch.Tensor, input: torch.Tensor, output_split_sizes: collections.abc.Sequence[typing.SupportsInt], input_split_sizes: collections.abc.Sequence[typing.SupportsInt], timeout: datetime.timedelta | None = None) -> c10d::Work
235
+
236
+ 跨进程组从所有进程中 alltoall(全交换)输入张量。
237
+
238
+ 有关详细信息,请参见 torch.distributed.all_to_all()。
239
+
240
+ barrier(self: torch._C._distributed_c10d.ProcessGroup, opts: torch._C._distributed_c10d.BarrierOptions = <torch._C._distributed_c10d.BarrierOptions object at 0x7f0162745ab0>) -> c10d::Work
241
+
242
+ 然后所有进程一起离开调用。
243
+
244
+ 有关详细信息,请参见 torch.distributed.barrier()。
245
+
246
+ barrier(self: torch._C._distributed_c10d.ProcessGroup, timeout: datetime.timedelta | None = None) -> c10d::Work
247
+
248
+ 然后所有进程一起离开调用。
249
+
250
+ 有关详细信息,请参见 torch.distributed.barrier()。
251
+
252
+ broadcast(self: torch._C._distributed_c10d.ProcessGroup, tensors: collections.abc.Sequence[torch.Tensor], opts: torch._C._distributed_c10d.BroadcastOptions = <torch._C._distributed_c10d.BroadcastOptions object at 0x7f0162b7afb0>) -> c10d::Work
253
+
254
+ 将张量广播到进程组中的所有进程。
255
+
256
+ 有关详细信息,请参见 torch.distributed.broadcast()。
257
+
258
+ broadcast(self: torch._C._distributed_c10d.ProcessGroup, tensor: torch.Tensor, root: typing.SupportsInt, timeout: datetime.timedelta | None = None) -> c10d::Work
259
+
260
+ 将张量广播到进程组中的所有进程。
261
+
262
+ 有关详细信息,请参见 torch.distributed.broadcast()。
263
+
264
+ gather(self: torch._C._distributed_c10d.ProcessGroup, output_tensors: collections.abc.Sequence[collections.abc.Sequence[torch.Tensor]], input_tensors: collections.abc.Sequence[torch.Tensor], opts: torch._C._distributed_c10d.GatherOptions = <torch._C._distributed_c10d.GatherOptions object at 0x7f0162c301f0>) -> c10d::Work
265
+
266
+ 跨进程组从所有进程中 gather(收集)输入张量。
267
+
268
+ 有关详细信息,请参见 torch.distributed.gather()。
269
+
270
+ gather(self: torch._C._distributed_c10d.ProcessGroup, output_tensors: collections.abc.Sequence[torch.Tensor], input_tensor: torch.Tensor, root: typing.SupportsInt, timeout: datetime.timedelta | None = None) -> c10d::Work
271
+
272
+ 跨进程组从所有进程中 gather(收集)输入张量。
273
+
274
+ 有关详细信息,请参见 torch.distributed.gather()。
275
+
276
+ 获取此进程组的存储。
277
+
278
+ 获取此进程组的描述
279
+
280
+ (获取此进程组的名称。它在集群中是唯一的)
281
+
282
+ 然后所有进程一起离开调用。
283
+
284
+ 有关详细信息,请参见 torch.distributed.monitored_barrier()。
285
+
286
+ 获取此进程组的名称。
287
+
288
+ 获取此进程组的 rank。
289
+
290
+ 从指定的 rank 接收张量。
291
+
292
+ 有关详细信息,请参见 torch.distributed.recv()。
293
+
294
+ 从任何来源接收张量。
295
+
296
+ 有关详细信息,请参见 torch.distributed.recv()。
297
+
298
+ reduce(self: torch._C._distributed_c10d.ProcessGroup, tensors: collections.abc.Sequence[torch.Tensor], opts: torch._C._distributed_c10d.ReduceOptions = <torch._C._distributed_c10d.ReduceOptions object at 0x7f0162bce3f0>) -> c10d::Work
299
+
300
+ 跨进程组中的所有进程归约提供的张量。
301
+
302
+ 有关详细信息,请参见 torch.distributed.reduce()。
303
+
304
+ reduce(self: torch._C._distributed_c10d.ProcessGroup, tensor: torch.Tensor, root: typing.SupportsInt, op: torch._C._distributed_c10d.ReduceOp = <RedOpType.SUM: 0>, timeout: datetime.timedelta | None = None) -> c10d::Work
305
+
306
+ 跨进程组中的所有进程归约提供的张量。
307
+
308
+ 有关详细信息,请参见 torch.distributed.reduce()。
309
+
310
+ reduce_scatter(self: torch._C._distributed_c10d.ProcessGroup, output_tensors: collections.abc.Sequence[torch.Tensor], input_tensors: collections.abc.Sequence[collections.abc.Sequence[torch.Tensor]], opts: torch._C._distributed_c10d.ReduceScatterOptions = <torch._C._distributed_c10d.ReduceScatterOptions object at 0x7f0162ee5cf0>) -> c10d::Work
311
+
312
+ 跨进程组从所有进程中归约并散射输入张量。
313
+
314
+ 有关详细信息,请参见 torch.distributed.reduce_scatter()。
315
+
316
+ reduce_scatter(self: torch._C._distributed_c10d.ProcessGroup, output: torch.Tensor, input: collections.abc.Sequence[torch.Tensor], op: torch._C._distributed_c10d.ReduceOp = <RedOpType.SUM: 0>, timeout: datetime.timedelta | None = None) -> c10d::Work
317
+
318
+ 跨进程组从所有进程中归约并散射输入张量。
319
+
320
+ 有关详细信息,请参见 torch.distributed.reduce_scatter()。
321
+
322
+ 跨进程组从所有进程中归约并散射输入张量。
323
+
324
+ 有关详细信息,请参见 torch.distributed.reduce_scatter()。
325
+
326
+ scatter(self: torch._C._distributed_c10d.ProcessGroup, output_tensors: collections.abc.Sequence[torch.Tensor], input_tensors: collections.abc.Sequence[collections.abc.Sequence[torch.Tensor]], opts: torch._C._distributed_c10d.ScatterOptions = <torch._C._distributed_c10d.ScatterOptions object at 0x7f0162b879f0>) -> c10d::Work
327
+
328
+ 跨进程组从所有进程中 scatter(散射)输入张量。
329
+
330
+ 有关详细信息,请参见 torch.distributed.scatter()。
331
+
332
+ scatter(self: torch._C._distributed_c10d.ProcessGroup, output_tensor: torch.Tensor, input_tensors: collections.abc.Sequence[torch.Tensor], root: typing.SupportsInt, timeout: datetime.timedelta | None = None) -> c10d::Work
333
+
334
+ 跨进程组从所有进程中 scatter(散射)输入张量。
335
+
336
+ 有关详细信息,请参见 torch.distributed.scatter()。
337
+
338
+ 将张量发送到指定的 rank。
339
+
340
+ 有关详细信息,请参见 torch.distributed.send()。
341
+
342
+ 为所有未来的操作设置默认超时时间。
343
+
344
+ 关闭进程组
345
+
346
+ 获取此进程组的大小。
347
+
348
+ 用于进程组工厂的协议。
349
+
350
+ 获取当前进程组。线程局部方法。
351
+
352
+ 当前进程组。
353
+
354
+ 使用给定的后端和选项创建一个新的进程组。这个组是独立的,不会被全局注册,因此无法通过标准的 torch.distributed.* API 使用。
355
+
356
+ backend (str) – 用于进程组的后端。
357
+
358
+ timeout (timedelta) – 集合通信操作的超时时间。
359
+
360
+ device (Union[str, device]) – 用于进程组的设备。
361
+
362
+ **kwargs (object) – 所有剩余的参数都将传递给后端构造函数。有关详细信息,请参见特定于后端的文档。
363
+
364
+ 进程组的上下文管理器。线程局部方法。
365
+
366
+ pg (ProcessGroup) – 要使用的进程组。
367
+
368
+ Generator[None, None, None]
369
+
370
+ 注册一个新的进程组后端。
371
+
372
+ name (str) – 后端的名称。
373
+
374
+ func (ProcessGroupFactory) – 用于创建进程组的函数。
375
+
376
+ ---
377
+
378
+ ## torch.distributed.fsdp.fully_shard#
379
+
380
+ **URL:** https://pytorch.org/docs/stable/distributed.fsdp.fully_shard.html
381
+
382
+ **目录:**
383
+ - torch.distributed.fsdp.fully_shard#
384
+ - PyTorch FSDP2 (fully_shard)#
385
+
386
+ 创建时间:2024年12月4日 | 最后更新时间:2025年6月16日
387
+
388
+ PyTorch FSDP2 (RFC) 提供了一个完全分片数据并行(FSDP)的实现,旨在实现高性能的 eager 模式,同时使用逐参数分片以提高可用性
389
+
390
+ 有关更多信息,请参阅 FSDP2 入门教程。
391
+
392
+ 如果您当前正在使用 FSDP1,请考虑使用我们的迁移指南迁移到 FSDP2。
393
+
394
+ fully_shard(model) 的用户契约如下
395
+
396
+ 对于模型初始化,fully_shard 原地将 model.parameters() 从普通的 torch.Tensor 转换为 DTensor。参数将根据设备网格移动到相应的设备上。
397
+
398
+ 在前向和反向传递之前,前向/反向前置钩子负责 all-gather 参数,并将 model.parameters() 从 DTensor 转换为普通的 torch.Tensor。
399
+
400
+ 在前向和反向传递之后,前向/反向后置钩子释放未分片的参数(不需要通信),并将 model.parameters() 从普通的 torch.Tensor 转换回 DTensor。
401
+
402
+ 对于优化器,它必须使用 DTensor 格式的 model.parameters() 进行初始化,并且优化器步骤应在 DTensor 参数上执行。
403
+
404
+ 请调用 model(input) 而不是 model.forward(input),以触发前向前置钩子来 all-gather 参数。要使 model.forward(input) 正常工作,用户必须显式调用 model.unshard() 或使用 register_fsdp_forward_method(model, "forward") 来注册该前向方法以进行挂钩。
405
+
406
+ fully_shard 将参数分组在一起以便进行单一的 all-gather。用户应以自下而上的方式应用 fully_shard。例如,在 Transformer 模型中,应将 fully_shard 应用于根模型之前的每一层。当应用于根模型时,fully_shard 会从每一层中排除 model.parameters(),并将剩余的参数(例如 embeddings、输出投影)分组到一个单一的 all-gather 组中。
407
+
408
+ type(model) 会原位与 FSDPModule 进行“联合”。例如,如果 model 最初是 nn.Linear 类型,则 fully_shard 会将 type(model) 从 nn.Linear 原位更改为 FSDPLinear。FSDPLinear 既是 nn.Linear 又是 FSDPModule 的实例。它保留了 nn.Linear 的所有方法,同时在 FSDPModule 下公开特定于 FSDP2 的 API,例如 reshard() 和 unshard()。
409
+
410
+ 参数的全限定名保持不变。如果我们调用 model.state_dict(),则在应用 fully_shard 之前和之后 FQN 是相同的。这是因为 fully_shard 没有包装模块,而只是将钩子注册到原始模块。
411
+
412
+ 与 PyTorch FSDP1 (FullyShardedDataParallel) 相比:
413
+
414
+ 与 FSDP1 的扁平参数分片相比,FSDP2 使用基于 DTensor 的 dim-0 逐参数分片以获得更简单的分片表示,同时保持相似的吞吐量性能。更具体地说,FSDP2 在数据并行工作进程之间对每个参数的 dim-0 进行分块(使用 torch.chunk(dim=0)),而 FSDP1 将一组张量展平、拼接后再进行分块,这使得推断每个工作进程上存在哪些数据以及重新分片为不同的并行策略变得复杂。逐参数分片提供了更直观的用户体验,放宽了对冻结参数的限制,并允许无通信的(分片)状态字典,而在 FSDP1 中这通常需要 all-gather。
415
+
416
+ FSDP2 实现了一种不同的内存管理方法来处理多流使用场景,避免了 torch.Tensor.record_stream。这确保了确定性和预期的内存使用情况,并且不需要像 FSDP1 的 limit_all_gathers=True 那样阻塞 CPU。
417
+
418
+ FSDP2 公开了用于手动控制预取和集合通信调度的 API,允许高级用户进行更多自定义。有关详细信息,请参见下文中的 FSDPModule 上的方法。
419
+
420
+ FSDP2 简化了一些 API 表面:例如,FSDP2 不直接支持完整状态字典。相反,用户可以使用 DTensor API(如 DTensor.full_tensor())或使用更高级别的 API(如 PyTorch 分布式检查点的分布式状态字典 API)自行将包含 DTensor 的分片状态字典重新分片为完整状态字典。此外,一些其他参数已被移除;有关详细信息,请参见此处。
421
+
422
+ 可以在模块上调用前端 API fully_shard:
423
+
424
+ 将完全分片数据并行(FSDP)应用于模块,其中 FSDP 将模块参数、梯度和优化器状态跨数据并行工作进程进行分片,以通过通信成本换取内存的节省。
425
+
426
+ 在初始化时,FSDP 根据 mesh 给定的数据并行工作进程对模块的参数进行分片。在前向传递之前,FSDP 跨数据并行工作进程 all-gather 已分片的参数,以获取用于前向计算的未分片参数。如果 reshard_after_forward 为 True,则 FSDP 在前向传递之后释放未分片的参数,并在计算梯度之前的反向传递中重新 all-gather 它们。在计算梯度之后,FSDP 释放未分片的参数,并跨数据并行工作进程对未分片的梯度进行 reduce-scatter(归约-散射)。
427
+
428
+ 此实现将在 dim-0 上分片的 DTensor 表示为分片参数,而未分片的参数将与模块上的原始参数相同(例如,如果原来是 torch.Tensor,则为 torch.Tensor)。模块上的前向前置钩子会 all-gather 参数,而模块上的前向钩子(如果需要)会释放它们。类似的后向钩子会 all-gather 参数,随后释放参数并 reduce-scatter 梯度。
429
+
430
+ 由于将多个张量组合在一起进行一次集合通信对于通信效率至关重要,因此此实现使这种组合成为一等公民。在模块上调用 fully_shard() 会构造一个组,其中包含 module.parameters() 中的参数,但之前在对子模块的调用中已分配给组的参数除外。这意味着应在模型上自下而上地调用 fully_shard()。每个组的参数将在一次集合通信中被 all-gather,并且其梯度将在一次集合通信中被 reduce-scatter。将模型划分为多个组(“逐层”)可以实现峰值内存节省以及通信与计算的重叠。用户通常不应该只在最顶层的根模块上调用 fully_shard()。
431
+
432
+ module (Union[nn.Module, List[nn.Module]) – 要使用 FSDP 进行分片并组合在一起以进行通信的模块或模块列表。
433
+
434
+ mesh (可选[DeviceMesh]) – 此数据并行网格定义了分片和设备。如果是 1D,则参数在 1D 网格上以 (Shard(0),) 放置进行完全分片(FSDP)。如果是 2D,则参数在第 1 维上进行分片,并在第 0 维上复制(HSDP),以 (Replicate(), Shard(0)) 进行放置。网格的设备类型给出了用于通信的设备类型;如果是 CUDA 或类 CUDA 设备类型,我们使用当前设备。
435
+
436
+ reshard_after_forward (可选[Union[bool, int]]) – 这控制了前向传递后的参数行为,并且可以在内存和通信之间进行权衡:如果为 True,则这会在前向传递之后重新分片参数,并在反向传递中重新 all-gather。如果为 False,则这会在前向传递之后将未分片的参数保留在内存中,并避免在反向传递中执行 all-gather。为了获得最佳性能,我们通常为根模块设置为 False,因为在反向传递开始时通常立即需要根模块。如果为 None,则对于非根模块设置为 True,对于根模块设置为 False。如果是整数,则表示前向传递后要重新分片到的 world size。它应该是 mesh 分片维度大小的非平凡因子(即不包括 1 和维度大小本身)。选择之一可以是节点内大小(例如 torch.cuda.device_count())。这允许在反向传递中以较小的 world size 执行 all-gather,代价是比设置为 True 占用更高的内存。在前向传递之后,注册到模块的参数取决于此设置:如果为 True,注册的参数是分片参数;如果为 False,则是未分片的参数;否则为重新分片到较小网格的参数。要修改前向和反向传递之间的参数,注册的参数必须是分片参数。对于 False 或整数,这可以通过 reshard() 手动重新分片来完成。
437
+
438
+ 这控制了前向传递后的参数行为,并且可以在内存和通信之间进行权衡:
439
+
440
+ 如果为 True,则这会在前向传递之后重新分片参数,并在反向传递中重新 all-gather。
441
+
442
+ 如果为 False,则这会在前向传递之后将未分片的参数保留在内存中,并避免在反向传递中执行 all-gather。为了获得最佳性能,我们通常为根模块设置为 False,因为在反向传递开始时通常立即需要根模块。
443
+
444
+ 如果为 None,则对于非根模块设置为 True,对于根模块设置为 False。
445
+
446
+ 如果是整数,则表示前向传递后要重新分片到的 world size。它应该是 mesh 分片维度大小的非平凡因子(即不包括 1 和维度大小本身)。选择之一可以是节点内大小(例如 torch.cuda.device_count())。这允许在反向传递中以较小的 world size 执行 all-gather,代价是比设置为 True 占用更高的内存。
447
+
448
+ 在前向传递之后,注册到模块的参数取决于此设置:如果为 True,注册的参数是分片参数;如果为 False,则是未分片的参数;否则为重新分片到较小网格的参数。要修改前向和反向传递之间的参数,注册的参数必须是分片参数。对于 False 或整数,这可以通过 reshard() 手动重新分片来完成。
449
+
450
+ shard_placement_fn (可选[Callable[[nn.Parameter], 可选[Shard]]]) – 此可调用对象可用于覆盖参数的分片放置,以便在 dim-0 以外的维度上对参数进行分片。如果此可调用对象返回了一个 Shard 放置(非 None),则 FSDP 将根据该放置进行分片(例如 Shard(1))。如果在非零维度上进行分片,我们目前要求均匀分片,即该维度上的张量维度大小必须能被 FSDP 分片 mesh 大小整除。
451
+
452
+ mp_policy (MixedPrecisionPolicy) – 这控制了混合精度策略,它为此模块提供参数/归约混合精度。有关详细信息,请参见 MixedPrecisionPolicy。
453
+
454
+ offload_policy (OffloadPolicy) – 这控制了卸载策略,它提供了参数/梯度/优化器状态的卸载。有关详细信息,请参见 OffloadPolicy 及其子类。
455
+
456
+ ignored_params (可选[set[nn.Parameter]]) – 可选(Set[nn.Parameter]): FSDP 将忽略的参数集。它们不会被分片,也不会在初始化期间移动到设备上,也不会在反向传递中对其梯度进行归约。
457
+
458
+ 应用了 FSDP 的模块(原位操作)。
459
+
460
+ 重新分片模块的参数,如果未分片的参数已分配则释放它们,并将分片参数注册到模块。此方法不是递归的。
461
+
462
+ hook (Callable[[torch.Tensor], None]) – 用户自定义的 all-reduce 钩子,预期签名为 hook(reduce_output: torch.Tensor) -> None,其中如果仅使用 FSDP,reduce_output 为 reduce-scatter 输出;如果使用原生 HSDP,则为 all-reduce 输出。
463
+
464
+ stream (可选[torch.cuda.Stream]) – 运行 all-reduce 钩子的流。仅在不使用原生 HSDP 时才应设置此项。如果使用原生 HSDP,该钩子将在原生 HSDP all-reduce 内部使用的定义的 all-reduce 流中运行。
465
+
466
+ 设置用于通过集合通信发送和接收数据的临时暂存缓冲区是否应使用 ProcessGroup 本身提供的自定义优化分配器进行分配(如果有)。这可能允许 ProcessGroup 更高效。例如,当使用 NCCL 时,这使其能够通过 SHARP(对于 NVLink 和/或 InfiniBand)利用零拷贝传输。
467
+
468
+ 这不能与 set_custom_all_gather() 或 set_custom_reduce_scatter() 一起使用,因为这些 API 允许对每次通信进行更细粒度的控制,而此方法无法确定它们的暂存缓冲区分配策略。
469
+
470
+ enable (bool) – 是否开启 ProcessGroup 分配。
471
+
472
+ 覆盖默认的 all_gather 通信行为,以更好地控制通信和内存使用。有关详细信息,请参见 Comm 和 ReduceScatter。
473
+
474
+ comm (AllGather) – 自定义 all-gather 通信。
475
+
476
+ 覆盖默认的 reduce_scatter 通信行为,以更好地控制通信和内存使用。有关详细信息,请参见 Comm 和 ReduceScatter。
477
+
478
+ comm (ReduceScatter) – 自定义 reduce_scatter 通信。
479
+
480
+ 设置是否要求底层集合通信原语专门使用“求和”类型的归约,即使这会带来单独的额外预缩放或后缩放操作代价。例如,需要这样做是因为 NCCL 目前仅支持此类集合通信的零拷贝传输。
481
+
482
+ 注意:对于 MTIA 设备,这始终是隐式启用的。
483
+
484
+ 注意:如果在 FSDP 设置下使用了 set_all_reduce_hook,调用者需要确保跨 FSDP 单元的自定义 all-reduce 也遵循此策略,因为 FSDP 无法再自动处理该情况。
485
+
486
+ enable (bool) – 是否仅使用 ReduceOp.SUM 进行通信。
487
+
488
+ 为梯度归约设置自定义除法因子。这可能会使用 NCCL 的 PreMulSum 使用自定义归约操作,允许在归约之前乘以该因子。
489
+
490
+ factor (float) – 自定义
491
+
492
+ ## 分布式通信包 - torch.distributed#
493
+
494
+ **URL:** https://pytorch.org/docs/stable/distributed.html
495
+
496
+ **目录:**
497
+ - 分布式通信包 - torch.distributed#
498
+ - 后端#
499
+ - PyTorch 自带的后端#
500
+ - 使用哪个后端?#
501
+ - 常用环境变量#
502
+ - 选择要使用的网络接口#
503
+ - 其他 NCCL 环境变量#
504
+ - 基础#
505
+ - 初始化#
506
+ - TCP 初始化#
507
+
508
+ 创建时间:2017 年 7 月 12 日 | 最后更新时间:2025 年 9 月 4 日
509
+
510
+ 有关与分布式训练相关的所有功能的简要介绍,请参阅 PyTorch Distributed 概览。
511
+
512
+ torch.distributed 支持四种内置后端,每种后端具有不同的功能。下表显示了每种后端在 CPU 或 GPU 上可用的功能。对于 NCCL,GPU 指的是 CUDA GPU,而对于 XCCL 指的是 XPU GPU。
513
+
514
+ 只有用于构建 PyTorch 的 MPI 实现支持 CUDA 时,MPI 才支持 CUDA。
515
+
516
+ PyTorch 分布式包支持 Linux(稳定版)、MacOS(稳定版)和 Windows(原型版)。在 Linux 上默认情况下,Gloo 和 NCCL 后端会被构建并包含在 PyTorch 分布式中(NCCL 仅在使用 CUDA 构建时包含)。MPI 是一个可选后端,只有在从源代码构建 PyTorch 时才能包含它。(例如,在安装了 MPI 的主机上构建 PyTorch。)
517
+
518
+ 从 PyTorch v1.8 开始,Windows 支持除 NCCL 外的所有集合通信后端,如果 init_process_group() 的 init_method 参数指向一个文件,它必须遵循以下模式:
519
+
520
+ 本地文件系统,init_method="file:///d:/tmp/some_file"
521
+
522
+ 共享文件系统,init_method="file://////{machine_name}/{share_folder_name}/some_file"
523
+
524
+ 与 Linux 平台上一样,您可以通过设置环境变量 MASTER_ADDR 和 MASTER_PORT 来启用 TcpStore。
525
+
526
+ 过去,我们经常被问到:“我应该使用哪个后端?”。
527
+
528
+ 在进行使用 CUDA GPU 的分布式训练时使用 NCCL 后端。
529
+
530
+ 在进行使用 XPU GPU 的分布式训练时使用 XCCL 后端。
531
+
532
+ 在进行使用 CPU 的分布式训练时使用 Gloo 后端。
533
+
534
+ 带有 InfiniBand 互连的 GPU 主机
535
+
536
+ 使用 NCCL,因为它是目前唯一支持 InfiniBand 和 GPUDirect 的后端。
537
+
538
+ 带有以太网互连的 GPU 主机
539
+
540
+ 使用 NCCL,因为它目前为分布式 GPU 训练提供了最佳性能,特别是对于多进程单节点或多节点分布式训练。如果您在使用 NCCL 时遇到任何问题,请使用 Gloo 作为备用选项。(请注意,目前 Gloo 在 GPU 上的运行速度比 NCCL 慢。)
541
+
542
+ 带有 InfiniBand 互连的 CPU 主机
543
+
544
+ 如果您的 InfiniBand 启用了 IP over IB,请使用 Gloo,否则请使用 MPI。我们计划在即将发布的版本中为 Gloo 添加 InfiniBand 支持。
545
+
546
+ 带有以太网互连的 CPU 主机
547
+
548
+ 使用 Gloo,除非您有特定原因要使用 MPI。
549
+
550
+ 默认情况下,NCCL 和 Gloo 后端都会尝试寻找合适的网络接口来使用。如果自动检测到的接口不正确,您可以使用以下环境变量(适用于各自的后端)进行覆盖:
551
+
552
+ NCCL_SOCKET_IFNAME,例如 export NCCL_SOCKET_IFNAME=eth0
553
+
554
+ GLOO_SOCKET_IFNAME,例如 export GLOO_SOCKET_IFNAME=eth0
555
+
556
+ 如果您使用的是 Gloo 后端,您可以通过用逗号分隔来指定多个接口,如下所示:export GLOO_SOCKET_IFNAME=eth0,eth1,eth2,eth3。后端将以轮询的方式在这些接口上分配操作。所有进程必须在此变量中指定相同数量的接口,这一点至关重要。
557
+
558
+ 调试 - 如果 NCCL 发生故障,您可以设置 NCCL_DEBUG=INFO 来打印明确的警告信息以及基本的 NCCL 初始化信息。
559
+
560
+ 您还可以使用 NCCL_DEBUG_SUBSYS 获取有关 NCCL 特定方面的更多详细信息。例如,NCCL_DEBUG_SUBSYS=COLL 将打印集合调用的日志,这在调试挂起(尤其是由集合类型或消息大小不匹配引起的挂起)时非常有用。如果拓扑检测失败,设置 NCCL_DEBUG_SUBSYS=GRAPH 来检查详细的检测结果,并在需要 NCCL 团队提供进一步帮助时将其保存为参考会很有帮助。
561
+
562
+ 性能调优 - NCCL 根据其拓扑检测执行自动调优,以节省用户的调优精力。在某些基于 socket 的系统上,用户仍然可以尝试调整 NCCL_SOCKET_NTHREADS 和 NCCL_NSOCKS_PERTHREAD 来增加 socket 网络带宽。NCCL 已为一些云提供商(如 AWS 或 GCP)预调优了这两个环境变量。
563
+
564
+ 有关 NCCL 环境变量的完整列表,请参阅 NVIDIA NCCL 的官方文档。
565
+
566
+ 你可以使用 torch.distributed.ProcessGroupNCCL.NCCLConfig 和 torch.distributed.ProcessGroupNCCL.Options 进一步 tune NCCL 通信器。在解释器中使用 help(例如 help(torch.distributed.ProcessGroupNCCL.NCCLConfig))来了解更多关于它们的信息。
567
+
568
+ torch.distributed 包为在一台或多台机器上运行的多个计算节点间的多进程并行提供了 PyTorch 支持和通信原语。torch.nn.parallel.DistributedDataParallel() 类在此功能的基础上构建,作为任何 PyTorch 模型的包装器,提供同步分布式训练。这与 Multiprocessing 包 - torch.multiprocessing 和 torch.nn.DataParallel() 提供的并行类型不同,它支持多台网络连接的机器,并且用户必须为每个进程显式启动主训练脚本的单独副本。
569
+
570
+ 在单机同步的情况下,torch.distributed 或 torch.nn.parallel.DistributedDataParallel() 包装器在数据并行方面可能仍然比其他方法(包括 torch.nn.DataParallel())具有优势:
571
+
572
+ 每个进程维护自己的优化器,并在每次迭代中执行完整的优化步骤。虽然这看起来可能是冗余的,因为梯度已经跨进程收集并求平均,因此对于每个进程来说都是相同的,但这意味着不需要参数广播步骤,从而减少了在节点之间传输张量所花费的时间。
573
+
574
+ 每个进程都包含一个独立的 Python 解释器,从而消除了从单个 Python 进程驱动多个执行线程、模型副本或 GPU 所带来的额外解释器开销和“GIL 抖动”。这对于大量使用 Python 运行时的模型(包括具有循环层或许多小组件的模型)尤为重要。
575
+
576
+ 在调用任何其他方法之前,需要使用 torch.distributed.init_process_group() 或 torch.distributed.device_mesh.init_device_mesh() 函数初始化此包。两者都会阻塞,直到所有进程都加入。
577
+
578
+ 初始化不是线程安全的。进程组的创建应该从单个线程执行,以防止跨 rank 的“UUID”分配不一致,并防止初始化期间可能导致挂起的竞争。
579
+
580
+ 如果分布式包可用,则返回 True。
581
+
582
+ 否则,torch.distributed 不会暴露任何其他 API。目前,torch.distributed 可在 Linux、MacOS 和 Windows 上使用。从源代码构建 PyTorch 时,设置 USE_DISTRIBUTED=1 以启用它。目前,Linux 和 Windows 的默认值为 USE_DISTRIBUTED=1,MacOS 的默认值为 USE_DISTRIBUTED=0。
583
+
584
+ 初始化默认分布式进程组。
585
+
586
+ 这也将初始化分布式包。
587
+
588
+ 显式指定 store、rank 和 world_size。
589
+
590
+ 指定 init_method(一个 URL 字符串),指示在哪里/如何发现对等节点。可选择指定 rank 和 world_size,或者在 URL 中编码所有必需的参数并省略它们。
591
+
592
+ 如果两者都未指定,则假定 init_method 为“env://”。
593
+
594
+ backend (str 或 Backend, 可选) – 要使用的后端。根据构建时的配置,有效值包括 mpi、gloo、nccl、ucc、xccl 或由第三方插件注册的后端。从 2.6 版本开始,如果未提供 backend,c10d 将使用为 device_id 关键字参数(如果提供)指示的设备类型注册的后端。目前已知的默认注册是:cuda 对应 nccl,cpu 对应 gloo,xpu 对应 xccl。如果未提供 backend 和 device_id,c10d 将检测运行时机器上的加速器,并使用为检测到的加速器(或 cpu)注册的后端。此字段可以作为小写字符串给出(例如,"gloo"),也可以通过 Backend 属性访问(例如,Backend.GLOO)。如果在每台机器上使用具有 nccl 后端的多个进程,则每个进程必须对其使用的每个 GPU 具有独占访问权,因为在进程之间共享 GPU 可能导致死锁或 NCCL 无效使用。ucc 后端是实验性的。可以通过 get_default_backend_for_device() 查询设备的默认 backend。
595
+
596
+ init_method (str, 可选) – 指定如何初始化进程组的 URL。如果没有指定 init_method 或 store,则默认为“env://”。与 store 互斥。
597
+
598
+ world_size (int, 可选) – 参与作业的进程数。如果指定了 store 则为必需。
599
+
600
+ rank (int, 可选) – 当前进程的 rank(它应该是 0 和 world_size-1 之间的数字)。如果指定了 store 则为必需。
601
+
602
+ store (Store, 可选) – 所有工作进程均可访问的键/值存储,用于交换连接/地址信息。与 init_method 互斥。
603
+
604
+ timeout (timedelta, 可选) – 针对进程组执行的操作的超时时间。NCCL 的默认值为 10 分钟,其他后端为 30 分钟。在此持续时间之后,集合通信将被异步中止,进程将崩溃。这样做是因为 CUDA 执行是异步的,并且由于失败的异步 NCCL 操作可能导致后续 CUDA 操作在损坏的数据上运行,因此继续执行用户代码不再安全。设置了 TORCH_NCCL_BLOCKING_WAIT 时,进程将阻塞并等待此超时。
605
+
606
+ group_name (str, 可选, 已弃用) – 组名。此参数被忽略
607
+
608
+ pg_options (ProcessGroupOptions, 可选) – 进程组选项,指定在构造特定进程组期间需要传入哪些附加选项。目前,我们支持的唯一选项是用于 nccl 后端的 ProcessGroupNCCL.Options,可以指定 is_high_priority_stream,以便在有计算内核等待时,nccl 后端可以选取高优先级的 CUDA 流。有关配置 nccl 的其他可用选项,请参阅 https://docs.nvidia.com/deeplearning/nccl/user-guide/docs/api/types.html#ncclconfig-t
609
+
610
+ device_id (torch.device | int, 可选) – 此进程将在其上运行的单个特定设备,允许进行特定于后端的优化。目前这有两个效果,仅在 NCCL 下:通信器立即形成(立即调用 ncclCommInit* 而不是正常的惰性调用),并且子组将在可能的情况下使用 ncclCommSplit 以避免组创建的不必要开销。如果您想及早了解 NCCL 初始化错误,也可以使用此字段。如果提供的是 int,则 API 假定将使用编译时的加速器类型。
611
+
612
+ 要启用 backend == Backend.MPI,需要在支持 MPI 的系统上从源代码构建 PyTorch。
613
+
614
+ 对多个后端的支持是实验性的。目前,当未指定后端时,将同时创建 gloo 和 nccl 后端。gloo 后端将用于具有 CPU 张量的集合通信,而 nccl 后端将用于具有 CUDA 张量的集合通信。可以通过传入格式为“<device_type>:<backend_name>,<device_type>:<backend_name>”的字符串(例如“cpu:gloo,cuda:custom_backend”)来指定自定义后端。
615
+
616
+ 根据 device_type、mesh_shape 和 mesh_dim_names 参数初始化 DeviceMesh。
617
+
618
+ 这将创建一个具有 n 维数组布局的 DeviceMesh,其中 n 是 mesh_shape 的长度。如果提供了 mesh_dim_names,则每个维度标记为 mesh_dim_names[i]。
619
+
620
+ init_device_mesh 遵循 SPMD 编程模型,这意味着相同的 PyTorch Python 程序在集群中的所有进程/rank 上运行。确保 mesh_shape(描述设备布局的 nD 数组的维度)在所有 rank 上都是相同的。不一致的 mesh_shape 可能会导致挂起。
621
+
622
+ 如果未找到进程组,init_device_mesh 将在后台初始化分布式通信所需的分布式进程组或组群。
623
+
624
+ device_type (str) – mesh 的设备类型。目前支持:“cpu”、“cuda/cuda-like”、“xpu”。不允许传入带有 GPU 索引的设备类型,例如“cuda:0”。
625
+
626
+ mesh_shape (Tuple[int]) – 定义描述设备布局的多维数组维度的元组。
627
+
628
+ mesh_dim_names (Tuple[str], 可选) – 分配给描述设备布局的多维数组每个维度的 mesh 维度名称元组。其长度必须与 mesh_shape 的长度相匹配。mesh_dim_names 中的每个字符串必须是唯一的。
629
+
630
+ backend_override (Dict[int | str, tuple[str, Options] | str | Options], 可选) – 对将为每个 mesh 维度创建的部分或全部 ProcessGroup 的覆盖。每个键可以是维度的索引或其名称(如果提供了 mesh_dim_names)。每个值可以是一个包含后端名称及其选项的元组,也可以只是这两个组件中的一个(在这种情况下,另一个将被设置为其默认值)。
631
+
632
+ 一个代表设备布局的 DeviceMesh 对象。
633
+
634
+ 检查默认进程组是否已初始化。
635
+
636
+ 检查 MPI 后端是否可用。
637
+
638
+ 检查 NCCL 后端是否可用。
639
+
640
+ 检查 Gloo 后端是否可用。
641
+
642
+ 检查 XCCL 后端是否可用。
643
+
644
+ 检查此进程是否由 torch.distributed.elastic(又名 torchelastic)启动。
645
+
646
+ TORCHELASTIC_RUN_ID 环境变量的存在被用作确定当前进程是否由 torchelastic 启动的代理。这是一个合理的代理,因为 TORCHELASTIC_RUN_ID 映射到 rendezvous id,该 id 始终为非空值,指示用于对等发现目的的作业 ID。
647
+
648
+ 返回给定设备的默认后端。
649
+
650
+ device (Union[str, torch.device]) – 要获取其默认后端的设备。
651
+
652
+ 给定设备的默认后端,以小写字符串形式返回。
653
+
654
+ 目前支持三种初始化方法:
655
+
656
+ 使用 TCP 初始化有两种方法,都需要一个所有进程都可以访问的网络地址和所需的 world_size。第一种方法需要指定一个属于 rank 0 进程的地址。这种初始化方法要求所有进程都手动指定 rank。
657
+
658
+ 请注意,最新的分布式包中不再支持多播地址。group_name 也已弃用。
659
+
660
+ 另一种初始化方法利用了一个组中所有机器都可以看到并共享的文件系统,以及所需的 world_size。URL 应以 file:// 开头,并包含共享文件系统上一个(位于现有目录中的)不存在文件的路径。如果文件不存在,文件系统初始化将自动创建该文件,但不会删除该文件。因此,您有责任确保在下一次对相同文件路径/名称调用 init_process_group() 之前清理该文件。
661
+
662
+ 请注意,最新的分布式包中不再支持自动 rank 分配,并且 group_name 也已弃用。
663
+
664
+ 此方法假定文件系统支持使用 fcntl 进行锁定 - 大多数本地系统和 NFS 都支持它。
665
+
666
+ 此方法将始终创建文件,并尽最大努力在程序结束时清理和删除该文件。换句话说,每次使用 file init 方法进行初始化时都需要一个全新的空文件,以便初始化成功。如果再次使用先前初始化使用过的文件(碰巧未被清理),这是意外行为,通常会导致死锁和失败。因此,即使此方法会尽最大努力清理文件,如果自动删除碰巧未成功,您有责任确保在训练结束时删除该文件,以防止下次再次重复使用相同的文件。如果您打算在同一文件名上多次调用 init_process_group(),这一点尤为重要。换句话说,如果文件未被删除/清理,而您再次对该文件调用 init_process_group(),预计会发生失败。这里的经验法则是,确保每次调用 init_process_group() 时文件都不存在或为空。
667
+
668
+ 此方法将从环境变量中读取配置,允许完全自定义获取信息的方式。要设置的变量有:
669
+
670
+ MASTER_PORT - 必填;必须是 rank 0 机器上的一个空闲端口
671
+
672
+ MASTER_ADDR - 必填(rank 0 除外);rank 0 节点的地址
673
+
674
+ WORLD_SIZE - 必填;可以在此处设置,或在调用 init 函数时设置
675
+
676
+ RANK - 必填;可以在此处设置,或在调用 init 函数时设置
677
+
678
+ rank 为 0 的机器将用于建立所有连接。
679
+
680
+ 这是默认方法,意味着不必指定 init_method(或者可以是 env://)。
681
+
682
+ TORCH_GLOO_LAZY_INIT - 按需建立连接,而不是使用全网状连接,这可以大大改善非 all2all 操作的初始化时间。
683
+
684
+ 一旦运行了 torch.distributed.init_process_group(),就可以使用以下函数。要检查进程组是否已经初始化,请使用 torch.distributed.is_initialized()。
685
+
686
+ 用于后端的类枚举类。
687
+
688
+ 可用后端:GLOO、NCCL、UCC、MPI、XCCL 和其他注册后端。
689
+
690
+ 此类的值为小写字符串,例如 "gloo"。它们可以作为属性访问,例如 Backend.NCCL。
691
+
692
+ 这个类可以被直接调用来解析字符串,例如,Backend(backend_str) 将检查 backend_str 是否有效,如果有效则返回解析后的小写字符串。它也接受大写字符串,例如,Backend("GLOO") 返回 "gloo"。
693
+
694
+ Backend.UNDEFINED 条目存在,但仅用作某些字段的初始值。用户既不应直接使用它,也不应假定它存在。
695
+
696
+ 使用给定的名称和实例化函数注册一个新的后端。
697
+
698
+ 第三方 ProcessGroup 扩展使用此类的注册新后端。
699
+
700
+ name (str) – ProcessGroup 扩展的后端名称。它应与 init_process_group() 中的名称相匹配。
701
+
702
+ func (函数) – 实例化后端的函数处理程序。该函数应在后端扩展中实现,并接受四个参数,包括 store、rank、world_size 和 timeout。
703
+
704
+ extended_api (bool, 可选) – 后端是否支持扩展参数结构。默认: False。如果设置为 True,后端将获得一个 c10d::DistributedBackendOptions 实例,以及一个由后端实现定义的进程组选项对象。
705
+
706
+ device (str 或 str 列表, 可选) – 此后端支持的设备类型,例如 “cpu”、“cuda” 等。如果为 None,则假定同时支持“cpu”和“cuda”
707
+
708
+ 这种对第三方后端的支持是实验性的,可能会发生变化。
709
+
710
+ 返回给定进程组的后端。
711
+
712
+ group (ProcessGroup, 可选) – 要在其上工作的进程组。默认是通用的主进程组。如果指定了另一个特定的组,则调用进程必须是该组的一部分。
713
+
714
+ 给定进程组的后端,以小写字符串形式返回。
715
+
716
+ 返回当前进程在提供的组中的 rank,否则为默认组。
717
+
718
+ Rank 是分配给分布式进程组内每个进程的唯一标识符。它们始终是范围从 0 到 world_size 的连续整数。
719
+
720
+ group (ProcessGroup, 可选) – 要在其上工作的进程组。如果为 None,则使用默认进程组。
721
+
722
+ 进程组的 rank,如果不属于该组,则为 -1
723
+
724
+ 返回当前进程组中的进程数。
725
+
726
+ group (ProcessGroup, 可选) – 要在其上工作的进程组。如果为 None,则使用默认进程组。
727
+
728
+ 进程组的 world size,如果不属于该组,则为 -1
729
+
730
+ 在退出时通过调用 destroy_process_group() 清理资源非常重要。
731
+
732
+ 最简单的模式是在训练脚本中不再需要通信的时候(通常在 main() 结尾附近),通过调用 destroy_process_group() 并为 group 参数传入默认值 None 来销毁每个进程组和后端。该调用应在每个训练器进程上执行一次,而不是在外部进程启动器级别执行。
733
+
734
+ 如果 pg 中的所有 rank 在超时时间内未调用 destroy_process_group(),尤其是当应用程序中有多个进程组时(例如用于 N-D 并行),在退出时可能会挂起。这是因为 ProcessGroupNCCL 的析构函数会调用 ncclCommAbort,而它必须被共同调用,但是如果由 python 的 GC 调用,调用 ProcessGroupNCCL 的析构函数的顺序是不确定的。调用 destroy_process_group() 有助于确保在所有 rank 之间以一致的顺序调用 ncclCommAbort,并避免在 ProcessGroupNCCL 的析构函数执行期间调用 ncclCommAbort。
735
+
736
+ destroy_process_group 也可用于销毁单个进程组。一个用例可能是容错训练,其中进程组可能在运行期间被销毁,然后初始化一个新的进程组。在这种情况下,_在_调用 destroy 之后且随后初始化之前,使用 torch.distributed 原语以外的一些手段同步训练器进程是至关重要的。由于实现这种同步很困难,这种行为目前是不受支持/未经过测试的,并且被认为是一个已知问题。如果这个用例阻碍了您,请提交 github issue 或 RFC。
737
+
738
+ 默认情况下,集合通信在默认组(也称为 world)上运行,并要求所有进程都进入分布式函数调用。然而,一些工作负载可以从更细粒度的通信中获益。这就是分布式组发挥作用的地方。new_group() 函数可用于创建包含所有进程任意子集的新组。它返回一个不透明的组句柄,可以作为 group 参数提供给所有集合通信(集合通信是在某些众所周知的编程模式中用于交换信息的分布式函数)。
739
+
740
+ 创建一个新的分布式组。
741
+
742
+ 此函数要求主组中的所有进程(即作为分布式作业一部分的所有进程)都进入此函数,即使它们不打算成为该组的成员。此外,在所有进程中,组应该以相同的顺序创建。
743
+
744
+ 安全的并发使用:当使用具有 NCCL 后端的多个进程组时,用户必须确保跨 rank 的集合通信具有全局一致的执行顺序。
745
+
746
+ 如果进程内的多个线程发出集合通信,则必须进行显式同步以确保一致的顺序。
747
+
748
+ 当使用 torch.distributed 通信 API 的异步变体时,将返回一个 work 对象,并且通信内核会在单独的 CUDA 流上排队,从而允许通信和计算重叠。一旦在一个进程组上发出了一个或多个异步操作,在使用另一个进程组之前,必须通过调用 work.wait() 将它们与其他 CUDA 流同步。
749
+
750
+ 有关更多详细信息,请参阅并发使用多个 NCCL 通信器 <https://docs.nvidia.com/deeplearning/nccl/user-guide/docs/usage/communicators.html#using-multiple-nccl-communicators-concurrently>。
751
+
752
+ ranks (list[int]) – 组成员 rank 的列表。如果为 None,将被设置为所有 rank。默认为 None。
753
+
754
+ timeout (timedelta, 可选) – 有关详细信息和默认值,请参见 init_process_group。
755
+
756
+ backend (str 或 Backend, 可选) – 要使用的后端。根据构建时的配置,有效值为 gloo 和 nccl。默认使用与全局组相同的后端。此字段应以小写字符串形式给出(例如 "gloo"),也可以通过 Backend 属性访问(例如 Backend.GLOO)。如果传入 None,将使用与默认进程组对应的后端。默认为 None。
757
+
758
+ pg_options (ProcessGroupOptions, 可选) – 进程组选项,指定在构造特定进程组期间需要传入哪些附加选项。即对于 nccl 后端,可以指定 is_high_priority_stream,以便进程组可以选取高优先级的 CUDA 流。有关配置 nccl 的其他可用选项,请参阅 https://docs.nvidia.com/deeplearning/nccl/user-guide/docs/api/types.html#ncclconfig-tuse_local_synchronization (bool, 可选): 在进程组创建结束时执行组本地屏障。这不同之处在于,非成员 rank 不需要调用 API 也不加入屏障。
759
+
760
+ group_desc (str, 可选) – 描述进程组的字符串。
761
+
762
+ device_id (torch.device, 可选) – 将此进程“绑定”到的单个特定设备,如果提供此字段,new_group 调用将尝试立即为该设备初始化通信后端。
763
+
764
+ 一个分布式组的句柄,可提供给集合调用,或者如果该 rank 不属于 ranks,则为 GroupMember.NON_GROUP_MEMBER。
765
+
766
+ 注意 use_local_synchronization 无法与 MPI 一起使用。
767
+
768
+ 注意 虽然在集群较大且进程组较小的情况下,use_local_synchronization=True 可以显著提高速度,但必须小心,因为它改变了集群行为,因为非成员 rank 不加入组的 barrier()。
769
+
770
+ 注意 当每个 rank 创建多个重叠的进程组时,use_local_synchronization=True 可能导致死锁。为避免这种情况,请确保所有 rank 遵循相同的全局创建顺序。
771
+
772
+ 将全局 rank 转换为组 rank。
773
+
774
+ global_rank 必须是该组的一部分,否则会引发 RuntimeError。
775
+
776
+ group (ProcessGroup) – 要查找相对 rank 的 ProcessGroup。
777
+
778
+ global_rank (int) – 要查询的全局 rank。
779
+
780
+ global_rank 相对于组的组 rank
781
+
782
+ 注意 在默认进程组上调用此函数返回原值
783
+
784
+ 将组 rank 转换为全局 rank。
785
+
786
+ group_rank 必须是该组的一部分,否则会引发 RuntimeError。
787
+
788
+ group (ProcessGroup) – 要从中查找全局 rank 的 ProcessGroup。
789
+
790
+ group_rank (int) – 要查询的组 rank。
791
+
792
+ group_rank 相对于组的全局 rank
793
+
794
+ 注意 在默认进程组上调用此函数返回原值
795
+
796
+ 获取与组关联的所有 rank。
797
+
798
+ group (可选[ProcessGroup]) – 要从中获取所有 rank 的 ProcessGroup。如果为 None,则使用默认进程组。
799
+
800
+ 按组 rank 排序的全局 rank 列表。
801
+
802
+ DeviceMesh 是管理进程组(或 NCCL 通信器)的更高级别的抽象。它允许用户轻松创建节点间和节点内的进程组,而无需担心如何为不同的子进程组正确设置 rank,并且它有助于轻松管理这些分布式进程组。init_device_mesh() 函数可用于创建新的 DeviceMesh,并使用一个描述设备拓扑的 mesh 形状。
803
+
804
+ DeviceMesh 代表一个设备网格,其中设备的布局可以表示为 n-d 维数组,而 n-d 维数组的每个值是默认进程组 rank 的全局 ID。
805
+
806
+ DeviceMesh 可用于设置跨集群的 N 维设备连接,并管理用于 N 维并行的 ProcessGroup。通信可以分别发生在 DeviceMesh 的每个维度上。DeviceMesh 尊重用户已经选择的设备(即如果用户在 DeviceMesh 初始化之前调用 torch.cuda.set_device),并且如果用户事先未设置设备,它将为当前进程选择/设置设备。请注意,手动选择设备必须在 DeviceMesh 初始化之前进行。
807
+
808
+ 与 DTensor API 一起使用时,DeviceMesh 也可以用作上下文管理器。
809
+
810
+ DeviceMesh 遵循 SPMD 编程模型,这意味着相同的 PyTorch Python 程序在集群中的所有进程/rank 上运行。因此,用户需要确保 mesh 数组(描述设备的布局)在所有 rank 上是相同的。不一致的 mesh 会导致静默挂起。
811
+
812
+ device_type (str) – mesh 的设备类型。目前支持:“cpu”、“cuda/cuda-like”。
813
+
814
+ mesh (ndarray) – 描述设备布局的多维数组或整数张量,其中 ID 是默认进程组的全局 ID。
815
+
816
+ 代表设备布局的 DeviceMesh 对象。
817
+
818
+ 下面的程序以 SPMD 方式在每个进程/rank 上运行。在这个示例中,我们有 2 台主机,每台主机有 4 个 GPU。对 mesh 第一维的规约将跨列 (0, 4), .. 和 (3, 7) 进行,对 mesh 第二维的规约将跨行 (0, 1, 2, 3) 和 (4, 5, 6, 7) 进行。
819
+
820
+ 从现有的 ProcessGroup 或现有 ProcessGroup 列表中构造具有 device_type 的 DeviceMesh。
821
+
822
+ 构造的设备 mesh 的维度数等于传入的组数。例如,如果传入单个进程组,则生成的 DeviceMesh 是一维 mesh。如果传入包含 2 个进程组的列表,则生成的 DeviceMesh 是二维 mesh。
823
+
824
+ 如果传入了多于一个组,则必须提供 mesh 和 mesh_dim_names 参数。传入的进程组的顺序决定了 mesh 的拓扑。例如,第一个进程组将是 DeviceMesh 的第 0 维。传入的 mesh 张量必须与传入的进程组数量具有相同的维数,并且 mesh 张量中的维度顺序必须与传入的进程组的顺序相匹配。
825
+
826
+ group (ProcessGroup 或 list[ProcessGroup]) – 现有的 ProcessGroup 或现有的 ProcessGroup 列表。
827
+
828
+ device_type (str) – mesh 的设备类型。目前支持:“cpu”、“cuda/cuda-like”。不允许传入带有 GPU 索引的设备类型,例如“cuda:0”。
829
+
830
+ mesh (torch.Tensor 或 ArrayLike, 可选) – 描述设备布局的多维数组或整数张量,其中 ID 是默认进程组的全局 ID。默认为 None。
831
+
832
+ mesh_dim_names (tuple[str], 可选) – 分配给描述设备布局的多维数组每个维度的 mesh 维度名称元组。其长度必须与 mesh_shape 的长度相匹配。mesh_dim_names 中的每个字符串必须是唯一的。默认为 None。
833
+
834
+ 代表设备布局的 DeviceMesh 对象。
835
+
836
+ 返回所有 mesh 维度的 ProcessGroup 列表。
837
+
838
+ 一个 ProcessGroup 对象列表。
839
+
840
+ list[torch.distributed.distributed_c10d.ProcessGroup]
841
+
842
+ 返回此 rank 相对于 mesh 所有维度的相对索引。如果此 rank 不是 mesh 的一部分,则返回 None。
843
+
844
+ 返回由 mesh_dim 指定的单个 ProcessGroup,或者如果未指定 mesh_dim 并且 DeviceMesh 是一维的,则返回 mesh 中唯一的 ProcessGroup。
845
+
846
+ mesh_dim (str/python:int, 可选) – 它可以是 mesh 维度的名称或索引
847
+
848
+ None。 (mesh 维度的。默认是) –
849
+
850
+ 一个 ProcessGroup 对象。
851
+
852
+ 返回 DeviceMesh 的给定 mesh_dim 的本地 rank。
853
+
854
+ mesh_dim (str/python:int, 可选) – 它可以是 mesh 维度的名称或索引
855
+
856
+ None。 (mesh 维度的。默认是) –
857
+
858
+ 表示本地 rank 的整数。
859
+
860
+ 下面的程序以 SPMD 方式在每个进程/rank 上运行。在这个示例中,我们有 2 台主机,每台主机有 4 个 GPU。在 rank 0、1、2、3 上调用 mesh_2d.get_local_rank(mesh_dim=0) 将返回 0。在 rank 4、5、6、7 上调用 mesh_2d.get_local_rank(mesh_dim=0) 将返回 1。在 rank 0、4 上调用 mesh_2d.get_local_rank(mesh_dim=1) 将返回 0。在 rank 1、5 上调用 mesh_2d.get_local_rank(mesh_dim=1) 将返回 1。在 rank 2、6 上调用 mesh_2d.get_local_rank(mesh_dim=1) 将返回 2。在 rank 3、7 上调用 mesh_2d.get_local_rank(mesh_dim=1) 将返回 3。
861
+
862
+ 返回当前的全局 rank。
863
+
864
+ 同步发送一个张量。
865
+
866
+ NCCL 后端不支持 tag。
867
+
868
+ tensor (Tensor) – 要发送的张量。
869
+
870
+ dst (int) – 全局进程组上的目标 rank(与 group 参数无关)。目标 rank 不应与当前进程的 rank 相同。
871
+
872
+ group (ProcessGroup, 可选) – 要在其上工作的进程组。如果为 None,则使用默认进程组。
873
+
874
+ tag (int, 可选) – 用于将发送与远程接收进行匹配的标签
875
+
876
+ group_dst (int, 可选) – 组上的目标 rank。不能同时指定 dst 和 group_dst。
877
+
878
+ 同步接收一个张量。
879
+
880
+ NCCL 后端不支持 tag。
881
+
882
+ tensor (Tensor) – 要用接收到的数据填充的张量。
883
+
884
+ src (int, 可选) – 全局进程组上的源 rank(与 group 参数无关)。如果未指定,将从任何进程接收。
885
+
886
+ group (ProcessGroup, 可选) – 要在其上工作的进程组。如果为 None,则使用默认进程组。
887
+
888
+ tag (int, 可选) – 用于将接收与远程发送进行匹配的标签
889
+
890
+ group_src (int, 可选) – 组上的目标 rank。不能同时指定 src 和 group_src。
891
+
892
+ 发送方 rank,如果不属于该组,则为 -1
893
+
894
+ 在使用时,isend() 和 irecv() 返回分布式请求对象。通常,此对象的类型是未指定的,因为永远不应该手动创建它们,但它们保证支持两个方法:
895
+
896
+ is_completed() - 如果操作已完成,则返回 True
897
+
898
+ wait() - 将阻塞进程直到操作完成。一旦返回,保证 is_completed() 返回 True。
899
+
900
+ 异步发送一个张量。
901
+
902
+ 在请求完成之前修改张量会导致未定义的行为。
903
+
904
+ NCCL 后端不支持 tag。
905
+
906
+ 与阻塞式的 send 不同,isend 允许 src == dst rank,即向自己发送。
907
+
908
+ tensor (Tensor) – 要发送的张量。
909
+
910
+ dst (int) – 全局进程组上的目标 rank(与 group 参数无关)
911
+
912
+ group (ProcessGroup, 可选) – 要在其上工作的进程组。如果为 None,则使用默认进程组。
913
+
914
+ tag (int, 可选) – 用于将发送与远程接收进行匹配的标签
915
+
916
+ group_dst (int, 可选) – 组上的目标 rank。不能同时指定 dst 和 group_dst
917
+
918
+ 分布式请求对象。如果不属于该组,则为 None
919
+
920
+ 异步接收一个张量。
921
+
922
+ NCCL 后端不支持 tag。
923
+
924
+ 与阻塞式的 recv 不同,irecv 允许 src == dst rank,即从自己接收。
925
+
926
+ tensor (Tensor) – 要用接收到的数据填充的张量。
927
+
928
+ src (int, 可选) – 全局进程组上的源 rank(与 group 参数无关)。如果未指定,将从任何进程接收。
929
+
930
+ group (ProcessGroup, 可选) – 要在其上工作的进程组。如果为 None,则使用默认进程组。
931
+
932
+ tag (int, 可选) – 用于将接收与远程发送进行匹配的标签
933
+
934
+ group_src (int, 可选) – 组上的目标 rank。不能同时指定 src 和 group_src。
935
+
936
+ 分布式请求对象。如果不属于该组,则为 None
937
+
938
+ 同步发送 object_list 中的可序列化对象。
939
+
940
+ 类似于 send(),但可以传入 Python 对象。请注意,object_list 中的所有对象都必须是可序列化的才能被发送。
941
+
942
+ object_list (List[Any]) – 要发送的输入对象列表。每个对象都必须是可序列化的。接收方必须提供同等大小的列表。
943
+
944
+ dst (int) – 要将 object_list 发送到的目标 rank。目标 rank 基于全局进程组(与 group 参数无关)
945
+
946
+ group (可选[ProcessGroup]) – (ProcessGroup, 可选): 要在其上工作的进程组。如果为 None,则使用默认进程组。默认为 None。
947
+
948
+ device (torch.device, 可选) – 如果不为 None,则对象将被序列化并转换为张量,在发送之前移动到该设备。默认为 None。
949
+
950
+ group_dst (int, 可选) – 组上的目标 rank。必须指定 dst 和 group_dst 中的一个,但不能同时指定两者
951
+
952
+ use_batch (bool, 可选) – 如果为 True,则使用批量 p2p 操作而不是常规发送操作。这避免了初始化 2 个 rank 的通信器,并使用现有的整个组通信器。有关用法和假设,请参见 batch_isend_irecv。默认为 False。
953
+
954
+ 对于基于 NCCL 的进程组,对象的内部张量表示必须在通信发生之前移动到 GPU 设备。在这种情况下,使用的设备由 torch.cuda.current_device() 给出,用户有责任通过 torch.cuda.set_device() 确保每个 rank 都有一个独立的 GPU。
955
+
956
+ 对象集合通信在性能和可伸缩性方面存在许多严重的限制。详情请参见对象集合通信。
957
+
958
+ send_object_list() 隐式使用 pickle 模块,众所周知这是不安全的。可能会构造恶意 pickle 数据,在反序列化期间执行任意代码。仅使用您信任的数据调用此函数。
959
+
960
+ 不支持使用 GPU 张量调用 send_object_list() 并且效率低下,因为这会导致 GPU -> CPU 传输(因为张量会被 pickle)。请考虑改用 send()。
961
+
962
+ 同步接收 object_list 中的可序列化对象。
963
+
964
+ 类似于 recv(),但可以接收 Python 对象。
965
+
966
+ object_list (List[Any]) – 要接收到的对象列表。必须提供与发送列表大小相等的列表。
967
+
968
+ src (int, 可选) – 要从其接收 object_list 的源 rank。源 rank 基于全局进程组(与 group 参数无关)。如果为 None,将从任何 rank 接收。默认为 None。
969
+
970
+ group (可选[ProcessGroup]) – (ProcessGroup, 可选): 要在其上工作的进程组。如果为 None,则使用默认进程组。默认为 None。
971
+
972
+ device (torch.device, 可选) – 如果不为 None,则在此设备上接收。默认为 None。
973
+
974
+ group_src (int, 可选) – 组上的目标 rank。不能同时指定 src 和 group_src。
975
+
976
+ use_batch (bool, 可选) – 如果为 True,则使用批量 p2p 操作而不是常规发送操作。这避免了初始化 2 个 rank 的通信器,并使用现有的整个组通信器。有关用法和假设,请参见 batch_isend_irecv。默认为 False。
977
+
978
+ 发送方 rank。如果不属于该组,则为 -1。如果 rank 属于该组,object_list 将包含来自 src rank 发送的对象。
979
+
980
+ 对于基于 NCCL 的进程组,对象的内部张量表示必须在通信发生之前移动到 GPU 设备。在这种情况下,使用的设备由 torch.cuda.current_device() 给出,用户有责任通过 torch.cuda.set_device() 确保每个 rank 都有一个独立的 GPU。
981
+
982
+ 对象集合通信在性能和可伸缩性方面存在许多严重的限制。详情请参见对象集合通信。
983
+
984
+ recv_object_list() 隐式使用 pickle 模块,众所周知这是不安全的。可能会构造恶意 pickle 数据,在反序列化期间执行任意代码。仅使用您信任的数据调用此函数。
985
+
986
+ 不支持使用 GPU 张量调用 recv_object_list() 并且效率低下,因为这会导致 GPU -> CPU 传输(因为张量会被 pickle)。请考虑改用 recv()。
987
+
988
+ 异步发送或接收一批张量,并返回请求列表。
989
+
990
+ 处理 p2p_op_list 中的每个操作并返回相应的请求。目前支持 NCCL、Gloo 和 UCC 后端。
991
+
992
+ p2p_op_list (list[torch.distributed.distributed_c10d.P2POp]) – 点对点操作列表(每个运算符的类型为 torch.distributed.P2POp)。列表中 isend/irecv 的顺序很重要,它需要与远程端对应的 isend/irecv 相匹配。
993
+
994
+ 通过调用 op_list 中的相应操作返回的分布式请求对象列表。
995
+
996
+ list[torch.distributed.distributed_c10d.Work]
997
+
998
+ 请注意,当此 API 与 NCCL PG 后端一起使用时,用户必须使用 torch.cuda.set_device 设置当前的 GPU 设备,否则将导致意外的挂起问题。
999
+
1000
+ 此外,如果此 API 是传递给 dist.P2POp 的组中的第一次集合调用,则该组的所有 rank 都必须参与此 API 调用;否则,行为是未定义的。如果此 API 调用不是组中的第一次集合调用,则允许仅涉及该组部分 rank 的批量 P2P 操作。
1001
+
1002
+ 用于为 batch_isend_irecv 构建点对点操作的类。
1003
+
1004
+ 这个类构建 P2P 操作的类型、通信缓冲区、
1005
+
1006
+ ## DistributedDataParallel#
1007
+
1008
+ **URL:** https://pytorch.org/docs/stable/generated/torch.nn.parallel.DistributedDataParallel.html
1009
+
1010
+ **目录:**
1011
+ - DistributedDataParallel#
1012
+
1013
+ 在模块级别实现基于 `torch.distributed` 的分布式数据并行。
1014
+
1015
+ 此容器通过同步每个模型副本的梯度来提供数据并行。要进行同步的设备由输入的 `process_group` 指定,默认为整个 world。请注意,DistributedDataParallel 不会在参与的 GPU 之间对输入进行分块或以其他方式进行分片;用户需自行定义如何执行此操作,例如通过使用 `DistributedSampler`。
1016
+
1017
+ 另见:基础知识和使用 nn.parallel.DistributedDataParallel 代替多进程或 nn.DataParallel。适用于 `torch.nn.DataParallel` 的相同输入约束。
1018
+
1019
+ 创建此类要求 `torch.distributed` 已经通过调用 `torch.distributed.init_process_group()` 完成初始化。
1020
+
1021
+ 事实证明,对于单节点多 GPU 数据并行训练,DistributedDataParallel 明显快于 `torch.nn.DataParallel`。
1022
+
1023
+ 要在具有 N 个 GPU 的主机上使用 DistributedDataParallel,您应该生成 N 个进程,确保每个进程独占地在从 0 到 N-1 的单个 GPU 上工作。这可以通过为每个进程设置 `CUDA_VISIBLE_DEVICES` 或调用以下 GPU API 来完成,
1024
+
1025
+ 或者调用统一的加速器 API,
1026
+
1027
+ 其中 i 的范围是 0 到 N-1。在每个进程中,您应参考以下内容来构造此模块:
1028
+
1029
+ 或者您可以使用最新的 API 进行初始化:
1030
+
1031
+ 为了在每个节点上生成多个进程,您可以使用 `torch.distributed.launch` 或 `torch.multiprocessing.spawn`。
1032
+
1033
+ 关于分布式训练的所有功能简介,请参阅 PyTorch 分布式概述。
1034
+
1035
+ DistributedDataParallel 可以与 `torch.distributed.optim.ZeroRedundancyOptimizer` 结合使用,以减少每个 rank 的优化器状态内存占用。有关更多详细信息,请参阅 ZeroRedundancyOptimizer 教程。
1036
+
1037
+ 在使用 GPU 时,nccl 后端是目前最快且强烈推荐的后端。这适用于单节点和多节点分布式训练。
1038
+
1039
+ 此模块还支持混合精度分布式训练。这意味着您的模型可以拥有不同类型的参数,例如 fp16 和 fp32 的混合类型,对这些混合类型参数的梯度归约也能正常工作。
1040
+
1041
+ 如果您在一个进程上使用 `torch.save` 对模块进行检查点保存,并在其他进程上使用 `torch.load` 进行恢复,请确保为每个进程正确配置了 `map_location`。如果没有 `map_location`,`torch.load` 会将模块恢复到保存该模块的设备上。
1042
+
1043
+ 当模型在 M 个节点上以 batch=N 进行训练时,如果损失是在批次中的各个实例间求和(而不是像通常那样求平均),则其梯度将比在单个节点上以 batch=M*N 训练的相同模型小 M 倍(因为不同节点之间的梯度是求平均的)。当您希望获得与本地训练在数学上等价的训练过程时,您应该考虑到这一点。但在大多数情况下,您可以将 DistributedDataParallel 包装的模型、DataParallel 包装的模型和单 GPU 上的普通模型同等对待(例如,在等效批量大小下使用相同的学习率)。
1044
+
1045
+ 参数永远不会在进程之间广播。该模块会对梯度执行 all-reduce 步骤,并假定它们将在所有进程中以相同的方式被优化器修改。缓冲区(例如 BatchNorm 统计信息)会在每次迭代中从 rank 0 进程的模块广播到系统中的所有其他副本。
1046
+
1047
+ 如果您将 DistributedDataParallel 与分布式 RPC 框架结合使用,应始终使用 `torch.distributed.autograd.backward()` 来计算梯度,并使用 `torch.distributed.optim.DistributedOptimizer` 来优化参数。
1048
+
1049
+ 目前,DistributedDataParallel 对使用 `torch.utils.checkpoint()` 的梯度检查点提供有限的支持。如果检查点操作使用 `use_reentrant=False`(推荐),DDP 将按预期工作,没有任何限制。但是,如果检查点操作使用 `use_reentrant=True`(默认值),DDP 只在模型中没有未使用的参数且每层最多只检查点一次时按预期工作(确保不要向 DDP 传递 `find_unused_parameters=True`)。我们目前不支持某层被多次检查点的情况,也不支持检查点模型中存在未使用参数的情况。
1050
+
1051
+ 为了让非 DDP 模型从 DDP 模型加载状态字典,需要在加载前应用 `consume_prefix_in_state_dict_if_present()` 以去除 DDP 状态字典中的前缀 “module.”。
1052
+
1053
+ 构造函数、forward 方法以及输出的微分(或此模块输出的函数)是分布式同步点。如果不同的进程可能执行不同的代码,请考虑到这一点。
1054
+
1055
+ 此模块假定在创建时所有参数都已注册到模型中。以后不应添加或删除任何参数。缓冲区也是如此。
1056
+
1057
+ 此模块假定每个分布式进程的模型中注册的所有参数顺序都是相同的。模块本身将按照模型注册参数的相反顺序进行梯度 allreduce。换句话说,用户有责任确保每个分布式进程具有完全相同的模型,从而具有完全相同的参数注册顺序。
1058
+
1059
+ 此模块允许具有非行优先连续步长的参数。例如,您的模型可能包含一些参数的 `torch.memory_format` 为 `torch.contiguous_format`,而另一些参数的格式为 `torch.channels_last`。但是,不同进程中对应的参数必须具有相同的步长。
1060
+
1061
+ 此模块不能与 `torch.autograd.grad()` 一起使用(即,只有在将梯度累加到参数的 `.grad` 属性中时才起作用)。
1062
+
1063
+ 如果您计划将此模块与 nccl 后端或(使用 Infiniband 的)gloo 后端一起使用,并与使用多个 worker 的 DataLoader 结合使用,请将多进程启动方法更改为 forkserver(仅限 Python 3)或 spawn。遗憾的是,Gloo(使用 Infiniband)和 NCCL2 不是 fork 安全的,如果不更改此设置,您很可能会遇到死锁。
1064
+
1065
+ 在用 DistributedDataParallel 包装模型后,绝不应尝试更改模型的参数。因为,当用 DistributedDataParallel 包装模型时,DistributedDataParallel 的构造函数会在构建时为模型本身的所有参数注册额外的梯度归约函数。如果之后更改模型的参数,梯度归约函数将不再匹配正确的参数集。
1066
+
1067
+ 将 DistributedDataParallel 与分布式 RPC 框架结合使用目前处于实验阶段,可能会发生变化。
1068
+
1069
+ module (Module) – 要被并行化的模块
1070
+
1071
+ device_ids (list of int 或 torch.device) – CUDA 设备。1) 对于单设备模块,device_ids 只能包含一个设备 id,代表该进程对应的输入模块所在的唯一 CUDA 设备。或者,device_ids 也可以是 None。2) 对于多设备模块和 CPU 模块,device_ids 必须是 None。当在上述两种情况下 device_ids 为 None 时,前向传播的输入数据和实际模块都必须放置在正确的设备上。(默认:None)
1072
+
1073
+ CUDA 设备。1) 对于单设备模块,device_ids 只能包含一个设备 id,代表该进程对应的输入模块所在的唯一 CUDA 设备。或者,device_ids 也可以是 None。2) 对于多设备模块和 CPU 模块,device_ids 必须是 None。
1074
+
1075
+ 当在上述两种情况下 device_ids 为 None 时,前向传播的输入数据和实际模块都必须放置在正确的设备上。(默认:None)
1076
+
1077
+ output_device (int 或 torch.device) – 单设备 CUDA 模块的输出设备位置。对于多设备模块和 CPU 模块,它必须是 None,由模块本身决定输出位置。(默认:单设备模块为 device_ids[0])
1078
+
1079
+ broadcast_buffers (bool) – 启用在 forward 函数开始时同步(广播)模块缓冲区的标志。(默认:True)
1080
+
1081
+ init_sync (bool) – 是否在初始化期间进行同步,以验证参数形状并广播参数和缓冲区。警告:如果将其设置为 False,用户需自行确保所有 rank 上的权重相同。(默认:True)
1082
+
1083
+ process_group – 用于分布式数据 all-reduction 的进程组。如果为 None,将使用由 `torch.distributed.init_process_group()` 创建的默认进程组。(默认:None)
1084
+
1085
+ bucket_cap_mb – DistributedDataParallel 会将参数分桶,以便每个桶的梯度归约可以潜在地与反向计算重叠。bucket_cap_mb 控制以 MebiBytes (MiB) 为单位的桶大小。如果为 None,将使用默认大小 25 MiB。(默认:None)
1086
+
1087
+ find_unused_parameters (bool) – 从被包装模块的 forward 函数返回值中包含的所有张量遍历 autograd 图。未作为此图的一部分接收梯度的参数将被预先标记为准备好进行归约。此外,在被包装模块的 forward 函数中可能已经使用过但未参与损失计算因而也不会接收梯度的参数,也会被预先标记为准备好进行归约。(默认:False)
1088
+
1089
+ check_reduction – 此参数已被弃用。
1090
+
1091
+ gradient_as_bucket_view (bool) – 当设置为 True 时,梯度将是指向 allreduce 通信桶不同偏移量的视图。这可以减少峰值内存使用量,节省的内存大小将等于总梯度大小。此外,它避免了梯度和 allreduce 通信桶之间的拷贝开销。当梯度是视图时,不能对梯度调用 `detach_()`。如果遇到此类错误,请参考 `torch/optim/optimizer.py` 中的 `zero_grad()` 函数作为解决方案进行修复。请注意,在第一次迭代后梯度将变成视图,因此应在第一次迭代后检查峰值内存的节省情况。
1092
+
1093
+ static_graph (bool) – 当设置为 True 时,DDP 会知道被训练的图是静态的。静态图意味着 1) 已使用和未使用的参数集合在整个训练循环中不会改变;在这种情况下,用户是否设置 `find_unused_parameters = True` 都无关紧要。2) 图的训练方式在整个训练循环中不会改变(意味着不存在取决于迭代的控制流)。当 `static_graph` 设置为 True 时,DDP 将支持过去无法支持的情况:1) 可重入的反向传播。2) 多次激活检查点。3) 模型包含未使用的参数时的激活检查点。4) 存在在 forward 函数之外的模型参数。5) 当存在未使用的参数时可能会提高性能,因为当 `static_graph` 设置为 True 时,DDP 不会在每次迭代中搜索图来检测未使用的参数。要检查是否可以将 `static_graph` 设置为 True,一种方法是在之前的模型训练结束时检查 ddp 日志数据,如果 `ddp_logging_data.get("can_set_static_graph") == True`,通常您也可以将 `static_graph` 设置为 True。示例:>>> model_DDP = torch.nn.parallel.DistributedDataParallel(model) >>> # Training loop >>> ... >>> ddp_logging_data = model_DDP._get_ddp_logging_data() >>> static_graph = ddp_logging_data.get("can_set_static_graph")
1094
+
1095
+ 当设置为 True 时,DDP 会知道被训练的图是静态的。静态图意味着 1) 已使用和未使用的参数集合在整个训练循环中不会改变;在这种情况下,用户是否设置 `find_unused_parameters = True` 都无关紧要。2) 图的训练方式在整个训练循环中不会改变(意味着不存在取决于迭代的控制流)。当 `static_graph` 设置为 True 时,DDP 将支持过去无法支持的情况:1) 可重入的反向传播。2) 多次激活检查点。3) 模型包含未使用的参数时的激活检查点。4) 存在在 forward 函数之外的模型参数。5) 当存在未使用的参数时可能会提高性能,因为当 `static_graph` 设置为 True 时,DDP 不会在每次迭代中搜索图来检测未使用的参数。要检查是否可以将 `static_graph` 设置为 True,一种方法是在之前的模型训练结束时检查 ddp 日志数据,如果 `ddp_logging_data.get("can_set_static_graph") == True`,通常您也可以将 `static_graph` 设置为 True。
1096
+
1097
+ delay_all_reduce_named_params (list of tuple of str 和 torch.nn.Parameter) – 一个命名的参数列表,当 `param_to_hook_all_reduce` 中指定的参数的梯度准备就绪时,这些参数的 all reduce 操作将被延迟。DDP 的其他参数不适用于此参数中指定的命名参数,因为这些命名参数将被 DDP 归约器忽略。
1098
+
1099
+ param_to_hook_all_reduce (torch.nn.Parameter) – 用于挂钩 `delay_all_reduce_named_params` 中指定的参数延迟 all reduce 操作的参数。
1100
+
1101
+ skip_all_reduce_unused_params – 当设置为 True 时,DDP 将跳过对未使用参数的归约。这要求在整个训练过程中,所有 rank 上的未使用参数保持一致。如果不满足此条件,可能会导致不同步并造成训练挂起。
1102
+
1103
+ module (Module) – 要被并行化的模块。
1104
+
1105
+ 用于在 DDP 中跨进程进行输入不均等训练的上下文管理器。
1106
+
1107
+ 此上下文管理器将跟踪已加入(joined)的 DDP 进程,并通过插入集合通信操作来“掩盖”前向和反向传播,以与非加入的 DDP 进程创建的操作相匹配。这将确保每个集合通信调用都有一个对应的由已加入的 DDP 进程发出的调用,从而防止在跨进程输入不均等进行训练时可能出现的挂起或错误。或者,如果标志 `throw_on_early_termination` 指定为 True,一旦某个 rank 的输入耗尽,所有训练器都会抛出错误,从而允许根据应用程序逻辑捕获和处理这些错误。
1108
+
1109
+ 一旦所有 DDP 进程都已加入,上下文管理器会将最后加入的进程对应的模型广播给所有进程,以确保模型在所有进程上都是相同的(这是由 DDP 保证的)。
1110
+
1111
+ 要使用此功能启用跨进程输入不均等的训练,只需将此上下文管理器包裹在您的训练循环外即可。不需要对模型或数据加载进行进一步的修改。
1112
+
1113
+ 如果此上下文管理器包裹的模型或训练循环具有额外的分布式集合通信操作,例如模型前向传播中的 `SyncBatchNorm`,则必须启用 `throw_on_early_termination` 标志。这是因为此上下文管理器无法感知非 DDP 的集合通信。当任何一个 rank 的输入耗尽时,此标志将导致所有 rank 抛出异常,从而允许在所有 rank 上捕获这些错误并从中恢复。
1114
+
1115
+ divide_by_initial_world_size (bool) – 如果为 True,将使用启动 DDP 训练时的初始 world_size 来除以梯度。如果为 False,将计算有效的 world_size(尚未耗尽其输入的 rank 数量),并在 allreduce 期间除以该数量。设置 `divide_by_initial_world_size=True` 可确保每个输入样本(包括不均等的输入)在对全局梯度的贡献权重上是相等的。这是通过始终使用初始 world_size 来除以梯度来实现的,即使我们遇到不均等的输入也是如此。如果将其设置为 False,我们将使用剩余的节点数来除以梯度。这确保了与较小 world_size 训练的等效性,尽管这意味着不均等的输入会对全局梯度贡献更多。通常,对于训练作业最后几个输入不均等的情况,您会希望将其设置为 True。在输入数量存在巨大差异的极端情况下,将其设置为 False 可能会提供更好的结果。
1116
+
1117
+ enable (bool) – 是否启用不均等输入检测。在您知道参与进程之间的输入是均等的情况下,可以传入 `enable=False` 来禁用。默认为 True。
1118
+
1119
+ throw_on_early_termination (bool) – 当至少一个 rank 的输入耗尽时,是抛出错误还是继续训练。如果为 True,将在第一个 rank 到达数据末尾时抛出异常。如果为 False,将以较小的有效 world size 继续训练,直到所有 rank 都加入。请注意,如果指定了此标志,则 `divide_by_initial_world_size` 标志将被忽略。默认为 False。
1120
+
1121
+ DDP 加入挂钩通过在前向和反向传播中镜像通信来实现输入不均等的训练。
1122
+
1123
+ kwargs (dict) – 包含任何用于在运行时修改加入挂钩行为的关键字参数的字典;所有共享相同加入上下文管理器的 Joinable 实例都会接收到相同的 kwargs 值。
1124
+
1125
+ 如果为 True,则梯度将除以启动 DDP 时使用的初始 world_size。如果为 False,梯度将除以有效的 world_size(即非加入进程的数量),这意味着不均等的输入会对全局梯度贡献更多。通常,如果不均等的程度较小,应将其设置为 True,但在极端情况下,可以将其设置为 False 以获得可能更好的结果。默认为 True。
1126
+
1127
+ 用于禁用 DDP 进程间梯度同步的上下文管理器。
1128
+
1129
+ 在此上下文中,梯度将累积在模块变量上,随后将在退出上下文后的第一次前向-反向传播中进行同步。
1130
+
1131
+ 前向传播应包含在上下文管理器内部,否则梯度仍将被同步。
1132
+
1133
+ 为用户自定义的跨多个 worker 的 DDP 梯度聚合注册通信挂钩。
1134
+
1135
+ 这个挂钩对于研究人员尝试新想法非常有用。例如,此挂钩可用于实现诸如 GossipGrad 和梯度压缩等几种算法,这些算法在运行分布式数据并行训练时涉及不同的参数同步通信策略。
1136
+
1137
+ state (object) – 传递给挂钩以在训练过程中维护任何状态信息。示例包括梯度压缩中的误差反馈、GossipGrad 中下一个要通信的节点等。它由每个 worker 在本地存储,并由该 worker 上的所有梯度张量共享。
1138
+
1139
+ 传递给挂钩以在训练过程中维护任何状态信息。示例包括梯度压缩中的误差反馈、GossipGrad 中下一个要通信的节点等。
1140
+
1141
+ 它由每个 worker 在本地存储,并由该 worker 上的所有梯度张量共享。
1142
+
1143
+ hook (Callable) – 具有以下签名的可调用对象:`hook(state: object, bucket: dist.GradBucket) -> torch.futures.Future[torch.Tensor]`:此函数在桶准备就绪时被调用。挂钩可以执行所需的任何处理,并返回一个指示任何异步工作(如 allreduce)完成的 Future。如果挂钩不执行任何通信,它仍必须返回一个已完成的 Future。Future 应包含梯度桶张量的新值。一旦桶准备就绪,c10d 归约器将调用此挂钩,并使用 Future 返回的张量将梯度复制到各个参数。请注意,future 的返回类型必须是单个张量。我们还提供了一个名为 `get_future` 的 API,用于检索与 `c10d.ProcessGroup.Work` 完成相关联的 Future。`get_future` 目前支持 NCCL,并且也支持 GLOO 和 MPI 上的大多数操作,但不支持点对点操作(send/recv)。
1144
+
1145
+ 具有以下签名的可调用对象:`hook(state: object, bucket: dist.GradBucket) -> torch.futures.Future[torch.Tensor]`:
1146
+
1147
+ 此函数在桶准备就绪时被调用。挂钩可以执行所需的任何处理,并返回一个指示任何异步工作(如 allreduce)完成的 Future。如果挂钩不执行任何通信,它仍必须返回一个已完成的 Future。Future 应包含梯度桶张量的新值。一旦桶准备就绪,c10d 归约器将调用此挂钩,并使用 Future 返回的张量将梯度复制到各个参数。请注意,future 的返回类型必须是单个张量。
1148
+
1149
+ 我们还提供了一个名为 `get_future` 的 API,用于检索与 `c10d.ProcessGroup.Work` 完成相关联的 Future。`get_future` 目前支持 NCCL,并且也支持 GLOO 和 MPI 上的大多数操作,但不支持点对点操作(send/recv)。
1150
+
1151
+ 梯度桶的张量不会预先除以 world_size。在进行诸如 allreduce 之类的操作时,用户有责任自行除以 world_size。
1152
+
1153
+ DDP 通信挂钩只能注册一次,并且应该在调用 backward 之前进行注册。
1154
+
1155
+ 挂钩返回的 Future 对象应包含一个与梯度桶内张量形状相同的单个张量。
1156
+
1157
+ `get_future` API 支持 NCCL,并部分支持 GLOO 和 MPI 后端(不支持诸如 send/recv 之类的点对点操作),并且将返回一个 `torch.futures.Future`。
1158
+
1159
+ 下面是一个返回相同张量的无操作挂钩示例。
1160
+
1161
+ 下面是一个并行 SGD 算法的示例,其中梯度在 allreduce 之前被编码,然后在 allreduce 之后被解码。
1162
+
1163
+ ---
1164
+
1165
+ ## DDP Communication Hooks#
1166
+
1167
+ **URL:** https://pytorch.org/docs/stable/ddp_comm_hooks.html
1168
+
1169
+ **目录:**
1170
+ - DDP Communication Hooks#
1171
+ - How to Use a Communication Hook?#
1172
+ - What Does a Communication Hook Operate On?#
1173
+ - 默认通信挂钩#
1174
+ - PowerSGD Communication Hook#
1175
+ - PowerSGD State#
1176
+ - PowerSGD Hooks#
1177
+ - Debugging Communication Hooks#
1178
+ - Checkpointing of Communication Hooks#
1179
+ - Acknowledgements#
1180
+
1181
+ Created On: Jun 06, 2025 | Last Updated On: Jun 06, 2025
1182
+
1183
+ DDP 通信挂钩是一个通用接口,用于通过覆盖 DistributedDataParallel 中原生的 allreduce 来控制如何跨 worker 通信梯度。系统提供了几个内置的通信挂钩,用户可以轻松应用这些挂钩中的任何一个来优化通信。此外,该挂钩接口还可以为更高级的用例支持用户自定义的通信策略。
1184
+
1185
+ 要使用通信挂钩,用户只需在训练循环之前让 DDP 模型注册该挂钩即可,如下所示。
1186
+
1187
+ torch.nn.parallel.DistributedDataParallel.register_comm_hook()
1188
+
1189
+ 通信挂钩提供了一种灵活的方式来进行梯度的 allreduce。因此,它主要在 allreduce 之前对每个副本上的梯度进行操作,这些梯度被分桶以增加通信和计算之间的重叠。特别地,`torch.distributed.GradBucket` 表示一桶需要进行 allreduce 的梯度张量。
1190
+
1191
+ 此类主要将扁平化的梯度张量(由 `buffer()` 返回)传递给 DDP 通信挂钩。该张量可以进一步分解为该桶内每个参数的独立张量列表(由 `get_per_parameter_tensors()` 返回),以应用逐层操作。
1192
+
1193
+ 由于桶在第一次迭代后会被重建,因此不应在训练开始时依赖这些索引。
1194
+
1195
+ 存储少数几层连续梯度的桶的索引。所有梯度都会被分桶。
1196
+
1197
+ 一个扁平化的 1D `torch.Tensor` 缓冲区,可以将其进一步分解为该桶内每个参数的独立张量列表。
1198
+
1199
+ 一个 `torch.Tensor` 列表。列表中的每个张量对应一个梯度。
1200
+
1201
+ 此桶是否是迭代中最后一个进行 allreduce 的桶。这也意味着此桶对应于前向传播中的最前面几层。
1202
+
1203
+ 使用输入张量缓冲区替换桶中的张量。
1204
+
1205
+ 一个 `torch.Tensor` 列表。列表中的每个张量对应一个模型参数。
1206
+
1207
+ 默认的通信挂钩是简单的无状态挂钩,因此 `register_comm_hook` 中的输入状态是进程组或 None。输入桶是一个 `torch.distributed.GradBucket` 对象。
1208
+
1209
+ 使用 GradBucket 张量调用 allreduce。
1210
+
1211
+ 一旦梯度张量在所有 worker 之间聚合完成,它的回调函数就会取平均值并返回结果。
1212
+
1213
+ 如果用户注册了此 DDP 通信挂钩,DDP 的结果预期将与未注册挂钩的情况相同。因此,这不会改变 DDP 的行为,用户可以将其作为参考,或修改此挂钩以记录有用的信息或在不影响 DDP 行为的情况下用于任何其他目的。
1214
+
1215
+ 通过将 GradBucket 强制转换为 `torch.float16` 并除以进程组大小来进行压缩。
1216
+
1217
+ 此 DDP 通信挂钩实现了一种简单的梯度压缩方法,将 GradBucket 张量转换为半精度浮点格式 (`torch.float16`),然后将其除以进程组大小。它对这些 float16 梯度张量进行 allreduce。一旦压缩的梯度张量被 allreduce,链式回调的解压操作会将其转换回输入数据类型(例如 float32)。
1218
+
1219
+ 警告:此 API 是实验性的,并且要求 NCCL 版本高于 2.9.6。
1220
+
1221
+ 此 DDP 通信挂钩实现了一种简单的梯度压缩方法,将 GradBucket 张量转换为半精度 Brain 浮点格式 (`torch.bfloat16`),然后将其除以进程组大小。它对这些 bfloat16 梯度张量进行 allreduce。一旦压缩的梯度张量被 allreduce,链式回调的解压操作会将其转换回输入数据类型(例如 float32)。
1222
+
1223
+ 此外,还提供了一个通信挂钩包装器,以支持将 `fp16_compress_hook()` 或 `bf16_compress_hook()` 作为包装器,这可以与其他通信挂钩结合使用。
1224
+
1225
+ 将输入张量转换为 `torch.float16`,将挂钩的结果转换回输入数据类型。
1226
+
1227
+ 此包装器将给定 DDP 通信挂钩的输入梯度张量转换为半精度浮点格式 (`torch.float16`),并将给定挂钩的结果张量转换回输入数据类型,例如 float32。因此,`fp16_compress_hook` 等同于 `fp16_compress_wrapper(allreduce_hook)`。
1228
+
1229
+ Callable[[Any, GradBucket], Future[Tensor]]
1230
+
1231
+ 警告:此 API 是实验性的,并且要求 NCCL 版本高于 2.9.6。
1232
+
1233
+ 此包装器将给定 DDP 通信挂钩的输入梯度张量转换为半精度 Brain 浮点格式 (`torch.bfloat16`),并将给定挂钩的结果张量转换回输入数据类型,例如 float32。
1234
+
1235
+ 因此,`bf16_compress_hook` 等同于 `bf16_compress_wrapper(allreduce_hook)`。
1236
+
1237
+ Callable[[Any, GradBucket], Future[Tensor]]
1238
+
1239
+ PowerSGD(Vogels 等人,NeurIPS 2019)是一种梯度压缩算法,可以提供非常高的压缩率并加速受限于带宽的分布式训练。此算法需要同时维护一些超参数和内部状态。因此,PowerSGD 通信挂钩是一个有状态挂钩,用户需要提供如下定义的状态对象。
1240
+
1241
+ 存储训练期间所有梯度的算法超参数和内部状态。
1242
+
1243
+ 特别是,`matrix_approximation_rank` 和 `start_powerSGD_iter` 是需要用户调整的主要超参数。为了性能,我们建议保持二进制超参数 `use_error_feedback` 和 `warm_start` 为开启状态。
1244
+
1245
+ `matrix_approximation_rank` 控制压缩低秩张量的大小,这决定了压缩率。秩越低,压缩越强。
1246
+
1247
+ 1.1. 如果 `matrix_approximation_rank` 太低,完整的模型质量将需要更多的训练步骤才能达到,或者永远无法达到,从而导致精度损失。
1248
+
1249
+ 1.2. `matrix_approximation_rank` 的增加会大幅增加压缩的计算成本,并且精度可能不会超过某个 `matrix_approximation_rank` 阈值而进一步提高。
1250
+
1251
+ 要调整 `matrix_approximation_rank`,我们建议从 1 开始并以 2 的倍数增加(如指数网格搜索,1、2、4...),直到达到令人满意的精度。通常只使用较小的值 1-4。对于某些 NLP 任务(如原论文的附录 D 所示),此值已增加到 32。
1252
+
1253
+ `start_powerSGD_iter` 将 PowerSGD 压缩推迟到步骤 `start_powerSGD_iter`,而在步骤 `start_powerSGD_iter` 之前运行原生的 allreduce。这种原生 allreduce + PowerSGD 的混合方案可以有效提高精度,即使使用相对较小的 `matrix_approximation_rank` 也是如此。这是因为训练阶段的开头通常对不准确的梯度非常敏感,过早压缩梯度可能会使训练迅速进入次优轨迹,从而对精度造成无法恢复的影响。
1254
+
1255
+ 要调整 `start_powerSGD_iter`,我们建议从总训练步数的 10% 开始,然后增加直到达到令人满意的精度。如果训练中包含预热阶段,`start_powerSGD_iter` 通常不应小于预热步数。
1256
+
1257
+ `min_compression_rate` 是压缩层时所需的最小压缩率。由于压缩会带来计算开销,只有在能够充分节省带宽的情况下,才值得压缩张量,即满足 `(num_rows + num_cols) * matrix_approximation_rank * min_compression_rate < num_rows * num_cols`。如果无法满足指定的压缩率阈值,则将直接对该张量进行 allreduce 而不进行压缩。
1258
+
1259
+ 一旦 PowerSGD 压缩开始,将每隔 `compression_stats_logging_frequency` 次迭代记录一次压缩统计数据。
1260
+
1261
+ `orthogonalization_epsilon` 可以是一个非常小的值(例如 1e-8),在正交化步骤中将其添加到每个归一化矩阵列中,以防止在任何列全为 0 时出现除以零错误。如果这已经被防止了(例如通过批归一化),为了准确性,建议将 epsilon 设置为 0。
1262
+
1263
+ `batch_tensors_with_same_shape` 控制是否在批处理操作中压缩和解压具有相同形状的张量,以实现更高的并行度。请注意,您还应该增加桶大小(即 DDP 构造函数中的 `bucket_cap_mb` 参数),以使更多具有相同形状的张量出现在同一个桶中,然而这可能会减少计算和通信之间的重叠,并由于堆叠相同形状的张量而增加内存占用。如果压缩/解压计算成为瓶颈,请设置为 True。
1264
+
1265
+ 如果启用了误差反馈或预热,DDP 中允许的 `start_powerSGD_iter` 的最小值为 2。这是因为 DDP 中还有另一个内部优化会在第 1 次迭代时重建桶,这可能会与在重建过程之前记忆的任何张量发生冲突。
1266
+
1267
+ PowerSGD 通常需要与模型梯度大小相同的额外内存来启用误差反馈,这可以补偿有偏的压缩通信并提高精度。
1268
+
1269
+ PowerSGD 挂钩可能与 Apex 自动混合精度包冲突。请改用 PyTorch 原生的自动混合精度包。
1270
+
1271
+ 实现 PowerSGD 算法。
1272
+
1273
+ 此 DDP 通信挂钩实现了论文中描述的 PowerSGD 梯度压缩算法。一旦梯度张量在所有 worker 之间聚合完成,此挂钩将按如下方式应用压缩:
1274
+
1275
+ 将输入的扁平化 1D 梯度张量视为每个参数的独立张量列表,并将所有张量分为两组:
1276
+
1277
+ 1.1 在 allreduce 之前应该被压缩的张量,因为压缩可以节省足够的带宽。
1278
+
1279
+ 1.2 其余张量将被直接进行 allreduce 而不进行压缩,包括所有的向量张量(例如偏置)。
1280
+
1281
+ 处理未压缩的张量:
1282
+
1283
+ 2.1. 为这些未压缩的张量分配连续的内存,并将所有未压缩的张量作为一个批次进行 allreduce,不进行压缩;
1284
+
1285
+ 2.2. 将单个未压缩的张量从连续内存复制回输入张量。
1286
+
1287
+ 处理应被 PowerSGD 压缩的张量:
1288
+
1289
+ 3.1. 对于每个张量 M,创建两个低秩张量 P 和 Q 用于分解 M,使得 M = PQ^T,其中 Q 从标准正态分布初始化并进行正交化;
1290
+
1291
+ 3.2. 计算 Ps 中的每个 P,等于 MQ;
1292
+
1293
+ 3.3. 将 Ps 作为一个批次进行 allreduce;
1294
+
1295
+ 3.4. 对 Ps 中的每个 P 进行正交化;
1296
+
1297
+ 3.5. 计算 Qs 中的每个 Q,近似等于 M^TP;
1298
+
1299
+ 3.6. 将 Qs 作为一个批次进行 allreduce;
1300
+
1301
+ 3.7. 计算所有压缩张量中的每个 M,近似等于 PQ^T。
1302
+
1303
+ 请注意,此通信挂钩在前 `state.start_powerSGD_iter` 次迭代中强制使用原生 allreduce。这不仅让用户能更好地控制速度和精度之间的权衡,也有助于为未来的通信挂钩开发者抽象掉 DDP 内部优化的某些复杂性。
1304
+
1305
+ state (PowerSGDState) – 用于配置压缩率并支持误差反馈、热启动等的状态信息。要调整压缩配置,主要需要调整 `matrix_approximation_rank`、`start_powerSGD_iter` 和 `min_compression_rate`。
1306
+
1307
+ bucket (dist.GradBucket) – 存储扁平化 1D 梯度张量(打包了多个按变量划分的张量)的桶。请注意,由于 DDP 通信挂钩仅支持单进程单设备模式,因此该桶中仅存储了刚好一个张量。
1308
+
1309
+ 通信的 Future 处理程序,它会原地更新梯度。
1310
+
1311
+ 实现简化的 PowerSGD 算法。
1312
+
1313
+ 此 DDP 通信挂钩实现了论文中描述的简化版 PowerSGD 梯度压缩算法。此变体不逐层压缩梯度,而是压缩打包了所有梯度的扁平化输入张量。因此,它比 `powerSGD_hook()` 更快,但通常会导致精度低得多,除非 `matrix_approximation_rank` 为 1。
1314
+
1315
+ 在这里增加 `matrix_approximation_rank` 不一定会提高精度,因为在不进行列/行对齐的情况下打包每个参数的张量可能会破坏低秩结构。因此,用户应始终优先考虑 `powerSGD_hook()`,只有当 `matrix_approximation_rank` 为 1 就能达到满意的精度时,才考虑此变体。
1316
+
1317
+ 一旦梯度张量在所有 worker 之间聚合完成,此挂钩将按如下方式应用压缩:
1318
+
1319
+ 将输入的扁平化 1D 梯度张量视为带 0 填充的方形张量 M;
1320
+
1321
+ 创建两个低秩张量 P 和 Q 用于分解 M,使得 M = PQ^T,其中 Q 从标准正态分布初始化并进行正交化;
1322
+
1323
+ 计算 P,等于 MQ;
1324
+
1325
+ 计算 Q,近似等于 M^TP;
1326
+
1327
+ 计算 M,近似等于 PQ^T。
1328
+
1329
+ 将输入张量截断为原始长度。
1330
+
1331
+ 请注意,此通信挂钩在前 `state.start_powerSGD_iter` 次迭代中强制使用原生 allreduce。这不仅让用户能更好地控制速度和精度之间的权衡,也有助于为未来的通信挂钩开发者抽象掉 DDP 内部优化的某些复杂性。
1332
+
1333
+ state (PowerSGDState) – 用于配置压缩率并支持误差反馈、热启动等的状态信息。要调整压缩配置,主要需要调整 `matrix_approximation_rank` 和 `start_powerSGD_iter`。
1334
+
1335
+ bucket (dist.GradBucket) – 存储扁平化 1D 梯度张量(打包了多个按变量划分的张量)的桶。请注意,由于 DDP 通信挂钩仅支持单进程单设备模式,因此该桶中仅存储了刚好一个张量。
1336
+
1337
+ 通信的 Future 处理程序,它会原地更新梯度。
1338
+
1339
+ 顾名思义,调试通信挂钩仅用于调试和性能优化的目的。
1340
+
1341
+ 调试通信挂钩不一定会输出正确的结果。
1342
+
1343
+ 返回一个包装了输入的 future,因此它是一个不会产生任何通信开销的无操作。
1344
+
1345
+ 此挂钩应仅用于 allreduce 优化的裕度分析,而不是用于常规的梯度同步。例如,如果注册此挂钩后只能观察到不到 10% 的训练时间加速,通常意味着 allreduce 对于此情况不是性能瓶颈。如果无法轻易获取 GPU 追踪或追踪分析由于 allreduce 和计算之间的重叠或各 rank 之间的不同步等某些因素而变得复杂时,这种检测会特别有用。
1346
+
1347
+ 有状态的通信挂钩可以作为模型检查点的一部分进行保存,以启用训练器重启。要使挂钩可序列化,应定义 `__setstate__` 和 `__getstate__`。
1348
+
1349
+ `__getstate__` 应从返回的字典中排除不可序列化的属性。
1350
+
1351
+ `__setstate__` 应正确初始化所提供状态中排除的不可序列化属性。
1352
+
1353
+ PowerSGDState 已经实现了 `__setstate__` 和 `__getstate__`,可以用作参考。
1354
+
1355
+ 返回一个将被序列化并保存的 `Dict[str, Any]`。
1356
+
1357
+ `process_group` 是不可序列化的,已从返回的状态中排除。
1358
+
1359
+ 接收提供的状态并设置到此 PowerSGDState 实例。
1360
+
1361
+ `process_group` 被设置为默认值。
1362
+
1363
+ 这里有一个保存和重新加载 PowerSGD 状态及挂钩的简单、端到端示例。
1364
+
1365
+ 非常感谢 PowerSGD 论文的作者 Thijs Vogels 对 PowerSGD 通信挂钩的代码审查,以及对比实验,这表明 PowerSGD 通信挂钩的性能与原论文中的实现不相上下。
1366
+
1367
+ ## 分布式检查点 - torch.distributed.checkpoint#
1368
+
1369
+ **URL:** https://pytorch.org/docs/stable/distributed.checkpoint.html
1370
+
1371
+ **目录:**
1372
+ - 分布式检查点 - torch.distributed.checkpoint#
1373
+ - 附加资源:#
1374
+
1375
+ 创建于:2022年11月16日 | 最后更新于:2025年9月4日
1376
+
1377
+ 分布式检查点 (DCP) 支持从多个 rank 并行加载和保存模型。它处理加载时的重新分片,这使得在一个集群拓扑中保存并在另一个集群拓扑中加载成为可能。
1378
+
1379
+ DCP 在几个重要方面与 `torch.save` 和 `torch.load` 不同:
1380
+
1381
+ 每次检查点它会生成多个文件,每个 rank 至少生成一个。
1382
+
1383
+ 它是原地操作的,这意味着模型应该首先分配其数据,然后 DCP 使用该存储空间。
1384
+
1385
+ 加载和保存检查点的入口点如下:
1386
+
1387
+ 分布式检查点 (DCP) 入门
1388
+
1389
+ 使用分布式检查点 (DCP) 进行异步保存
1390
+
1391
+ TorchTitan 检查点文档
1392
+
1393
+ TorchTitan DCP 实现
1394
+
1395
+ 用于异步检查点类型的枚举。
1396
+
1397
+ 此类包含用于暂存和上传完成的 Future 对象。它由 `async_save()` 返回。`staging_completion` 是一个 Future,指示 `state_dict` 的本地副本何时完成。`upload_completion` 是一个 Future,指示检查点何时完成保存。
1398
+
1399
+ 以 SPMD(单程序多数据)风格保存分布式模型。
1400
+
1401
+ 此函数与 `torch.save()` 不同,因为它处理 `ShardedTensor` 和 `DTensor`,方式是让每个 rank 仅保存其本地分片。
1402
+
1403
+ 对于每个有状态对象(同时具有 `state_dict` 和 `load_state_dict`),保存时将在序列化之前调用 `state_dict`。
1404
+
1405
+ 对于保存的 `state_dict`,不保证跨 PyTorch 版本的向后兼容性。
1406
+
1407
+ 如果使用 `process_group` 参数,请确保只有其 rank 调用 `save_state_dict`,并且 `state_dict` 中的所有数据都属于该进程组。
1408
+
1409
+ 当为 FSDP 的 `ShardingStrategy.HYBRID_SHARD` 保存检查点时,`shard_group` 中应只有一个在调用 `save_state_dict`,并且需要传入相应的进程组。
1410
+
1411
+ 本地进程中的 `state_dict`。
1412
+
1413
+ state_dict (Dict[str, Any]) – 要保存的 state_dict。
1414
+
1415
+ checkpoint_id (Union[str, os.PathLike, None]) – 此检查点实例的 ID。`checkpoint_id` 的含义取决于存储方式。它可以是文件夹或文件的路径。如果存储是键值存储,它也可以是一个键。(默认: None)
1416
+
1417
+ storage_writer (可选[StorageWriter]) – 用于执行写入操作的 `StorageWriter` 实例。如果未指定,DCP 将根据 `checkpoint_id` 自动推断写入器。如果 `checkpoint_id` 也为 None,则会引发异常。(默认: None)
1418
+
1419
+ planner (可选[SavePlanner]) – `SavePlanner` 的实例。如果未指定,将使用默认规划器。(默认: None)
1420
+
1421
+ process_group (可选[ProcessGroup]) – 用于跨 rank 同步的进程组。(默认: None)
1422
+
1423
+ no_dist (bool) – 如果为 True,此函数将假定意图是在单个 rank/进程上加载检查点。(默认: False)
1424
+
1425
+ use_collectives (bool) – 如果为 False,此函数将假定意图是在不使用跨 rank 同步的情况下保存检查点。(默认: True) 此配置是实验性的,应谨慎使用。它会更改保存的检查点的格式,并且可能无法向后兼容。
1426
+
1427
+ 保存的检查点的元数据对象。
1428
+
1429
+ `save_state_dict` 使用集合通信来协调跨 rank 的写入。对于基于 NCCL 的进程组,对象的内部张量表示必须在通信发生之前移动到 GPU 设备。在这种情况下,所使用的设备由 `torch.cuda.current_device()` 给出,用户有责任通过 `torch.cuda.set_device()` 确保已设置此项,以便每个 rank 都有一个独立的 GPU。
1430
+
1431
+ `save` 的异步版本。此代码首先将 `state_dict` 暂存到暂存存储(默认为 CPU 内存),然后在一个单独的线程中调用保存。
1432
+
1433
+ 此功能是实验性的,可能会发生变化。在保存最后一个检查点后必须调用 CLOSE。
1434
+
1435
+ state_dict (Dict[str, Any]) – 要保存的 state_dict。
1436
+
1437
+ checkpoint_id (Union[str, os.PathLike, None]) – 此检查点实例的 ID。`checkpoint_id` 的含义取决于存储方式。它可以是文件夹或文件的路径。如果存储是键值存储,它也可以是一个键。(默认: None)
1438
+
1439
+ storage_writer (可选[StorageWriter]) – 用于执行“暂存”和“保存”的 `StorageWriter` 实例。如果未指定,DCP 将根据 `checkpoint_id` 自动推断写入器。如果 `checkpoint_id` 也为 None,则会引发异常。(默认: None)
1440
+
1441
+ planner (可选[SavePlanner]) – `SavePlanner` 的实例。如果未指定,将使用默认规划器。(默认: None)
1442
+
1443
+ process_group (可选[ProcessGroup]) – 用于跨 rank 同步的进程组。(默认: None)
1444
+
1445
+ async_checkpointer_type (AsyncCheckpointerType) – 是在单独的线程还是进程中执行检查点操作 (默认: AsyncCheckpointerType.THREAD)
1446
+
1447
+ async_stager (AsyncStager) – 提供暂存实现。如果 `storage_writer` 实现了 `AsyncStager` 并且提供了 `async_stager`,则将使用 `async_stager` 进行暂存
1448
+
1449
+ no_dist (bool) – 如果为 True,此函数将假定意图是在单个 rank/进程上保存检查点。(默认: False)
1450
+
1451
+ use_collectives (bool) – 如果为 False,在没有 rank 协调的情况下保存检查点。(默认: True) 此配置是实验性的,应谨慎使用。它会更改保存的检查点的格式,并且可能无法向后兼容。
1452
+
1453
+ 包含来自保存操作生成的 Metadata 对象的 Future。
1454
+
1455
+ 此方法已弃用。请切换到 ‘save’。
1456
+
1457
+ 以 SPMD 风格将检查点加载到分布式 state dict 中。
1458
+
1459
+ 每个 rank 在提供给此 API 的 `state_dict` 中必须具有相同的键。不匹配的键可能会导致挂起或错误。如果不确定,可以使用 `utils._assert_same_keys` API 进行检查(但可能会产生通信成本)。
1460
+
1461
+ 每个 rank 将尝试读取最少量的必要数据,以满足所请求的 `state_dict`。当加载 `ShardedTensor` 或 `DTensor` 实例时,每个 rank 仅读取其本地分片的数据。
1462
+
1463
+ 对于每个有状态对象(同时具有 `state_dict` 和 `load_state_dict`),加载将首先在尝试反序列化之前调用 `state_dict`,并在反序列化完成后调用 `load_state_dict`。对于每个非有状态对象,加载将对该对象进行反序列化,然后用反序列化后的对象替换 `state_dict` 中的该对象。
1464
+
1465
+ 在调用此函数之前,`state_dict` 中的所有张量都必须分配到其目标设备上。
1466
+
1467
+ 所有非张量数据都使用 `torch.load()` 加载,并在 `state_dict` 中进行原位修改。
1468
+
1469
+ 用户必须在根模块上调用 `load_state_dict`,以确保加载后处理和非张量数据的正确传播。
1470
+
1471
+ state_dict (Dict[str, Any]) – 要将检查点加载到的 state_dict。
1472
+
1473
+ checkpoint_id (Union[str, os.PathLike, None]) – 此检查点实例的 ID。`checkpoint_id` 的含义取决于存储方式。它可以是文件夹或文件的路径。如果存储是键值存储,它也可以是一个键。(默认: None)
1474
+
1475
+ storage_reader (可选[StorageReader]) – 用于执行读取操作的 `StorageReader` 实例。如果未指定,DCP 将根据 `checkpoint_id` 自动推断读取器。如果 `checkpoint_id` 也为 None,则会引发异常。(默认: None)
1476
+
1477
+ planner (可选[LoadPlanner]) – `LoadPlanner` 的实例。如果未指定,将使用默认规划器。(默认: None)
1478
+
1479
+ process_group (可选[ProcessGroup]) – 用于跨 rank 同步的进程组。(默认: None)
1480
+
1481
+ no_dist (bool) – 如果为 True,此函数将假定意图是在不使用跨 rank 同步的情况下加载检查点。(默认: False)
1482
+
1483
+ `load_state_dict` 使用集合通信来协调跨 rank 的读取。对于基于 NCCL 的进程组,对象的内部张量表示必须在通信发生之前移动到 GPU 设备。在这种情况下,所使用的设备由 `torch.cuda.current_device()` 给出,用户有责任通过 `torch.cuda.set_device()` 确保已设置此项,以便每个 rank 都有一个独立的 GPU。
1484
+
1485
+ 此方法已弃用。请切换到 ‘load’。
1486
+
1487
+ 以下模块对于进一步自定义用于异步检查点的暂存机制 也很有用:
1488
+
1489
+ 此协议旨在为 `dcp.async_save` 提供自定义和可扩展性,允许用户自定义在并行执行常规 `dcp.save` 路径之前如何暂存数据。预期的操作顺序(具体在 `torch.distributed.state_dict_saver.async_save` 中定义)如下:
1490
+
1491
+ 此调用给 `AsyncStager` 提供“暂存” `state_dict` 的机会。在此上下文中,暂存的期望和目的是创建 state dict 的“训练安全”表示,这意味着在暂存完成后对模块数据的任何更新都不应反映在此方法返回的 state dict 中。例如,在默认情况下,会在 CPU RAM 上创建整个 state dict 的副本并在此处返回,允许用户继续训练而不会冒正在序列化的数据被更改的风险。
1492
+
1493
+ 用于序列化 `state_dict` 并将其写入存储。
1494
+
1495
+ 序列化线程启动之前以及从 `dcp.async_save` 返回之前。如果将其设置为 False,则假定用户已经定义了一个自定义同步点,目的是为了进一步优化训练循环中的保存延迟(例如,将暂存与前向/后向传递重叠),并且用户有责任在适当的时间调用 `AsyncStager.synchronize_staging`。
1496
+
1497
+ 清理暂存器使用的所有资源。
1498
+
1499
+ 执行暂存后是否同步。
1500
+
1501
+ 返回 `state_dict` 的“暂存”副本。对暂存副本的期望是,它不受在暂存调用完成后发生的任何更新的影响。
1502
+
1503
+ Union[Future[dict[str, Union[~StatefulT, Any]]], dict[str, Union[~StatefulT, Any]]]
1504
+
1505
+ 如果暂存以某种方式是异步的,则应调用此方法以确保暂存完成,并且可以安全地开始修改原始的 `state_dict`。
1506
+
1507
+ `DefaultStager` 提供了一个功能齐全的暂存实现,它结合了多种优化技术以高效地准备检查点。
1508
+
1509
+ 暂存过程的工作原理如下: 1. 提交状态字典以进行暂存(同步或异步) 2. 将张量从 GPU 复制到优化的 CPU 存储中 3. 如果使用非阻塞复制,则同步 CUDA 操作 4. 通过 Future 返回或提供已暂存的状态字典
1510
+
1511
+ # 同步暂存 stager = DefaultStager(StagingOptions(use_async_staging=False)) staged_dict = stager.stage(state_dict) stager.close()
1512
+
1513
+ # 异步暂存 stager = DefaultStager(StagingOptions(use_async_staging=True)) future = stager.stage(state_dict) # …… 执行其他工作 …… staged_dict = future.result() stager.close()
1514
+
1515
+ # 上下文管理器模式(推荐) stager = DefaultStager(config) with stager: result = stager.stage(state_dict)
1516
+
1517
+ 当模型计算可以与暂存操作重叠时,异步暂存可提供最佳性能。
1518
+
1519
+ 固定内存可提高 CPU-GPU 传输速度,但会使用更多的内存。
1520
+
1521
+ 共享内存允许与检查点进程进行高效的 IPC(进程间通信)。
1522
+
1523
+ 非阻塞复制减少了内存传输期间的 GPU 空闲时间。
1524
+
1525
+ `DefaultStager` 不是线程安全的。每个线程应使用自己的实例,或者应提供外部同步。
1526
+
1527
+ 清理 `DefaultStager` 使用的所有资源。关闭用于异步暂存操作的 `ThreadPoolExecutor`,并清理底层 `StateDictStager` 的缓存存储。当不再需要暂存器时应调用此方法以防止资源泄漏,特别是在长时间运行的应用程序中。调用 `close()` 后,不应将其用于进一步的暂存操作。
1528
+
1529
+ stager = DefaultStager(StagingOptions(use_async_staging=True)) future = stager.stage(state_dict) result = future.result() stager.close() # 清理所有资源
1530
+
1531
+ 此函数负责暂存 `state_dict`。有关暂存的更多详细信息,请参见类文档字符串。如果 `use_async_staging` 为 True,它将返回一个在暂存完成时会被履行的 Future 对象。如果 `use_async_staging` 为 False,它将返回完全暂存的 `state_dict`。
1532
+
1533
+ state_dict (STATE_DICT_TYPE) – 要暂存的 state_dict。
1534
+
1535
+ Union[dict[str, Union[~StatefulT, Any]], Future[dict[str, Union[~StatefulT, Any]]]]
1536
+
1537
+ 当 `use_async_staging` 为 True 时,此方法将等待直到暂存完成。如果 `use_async_staging` 为 False,则此方法为无操作。
1538
+
1539
+ 配置检查点暂存行为的选项。
1540
+
1541
+ use_pinned_memory (bool) – 启用固定内存分配以实现更快的 CPU-GPU 传输。需要 CUDA 可用。 默认: True
1542
+
1543
+ use_shared_memory (bool) – 为多进程场景启用共享内存。当多个进程需要访问相同的暂存数据时非常有用。 默认: True
1544
+
1545
+ use_async_staging (bool) – 使用后台线程池启用异步暂存。允许将计算与暂存操作重叠。需要 CUDA。 默认: True
1546
+
1547
+ use_non_blocking_copy (bool) – 使用带有流同步的非阻塞设备内存复制。通过允许 CPU 工作在 GPU 传输期间继续进行来提高性能。 默认: True
1548
+
1549
+ 如果 CUDA 不可用,依赖于 CUDA 的功能将引发异常。
1550
+
1551
+ `AsyncStager` 的一种实现,它在 CPU RAM 上暂存 `state_dict` 并阻塞直到复制完成。此实现还提供了一个选项,以使用固定内存优化暂存延迟。
1552
+
1553
+ 注意:在这种情况下 `synchronize_staging` 是一个无操作的方法。
1554
+
1555
+ 返回 CPU 上的 `state_dict` 副本。
1556
+
1557
+ dict[str, Union[~StatefulT, Any]]
1558
+
1559
+ 无操作函数,因为暂存是阻塞的。
1560
+
1561
+ 除了上述入口点之外,如下所述的有状态对象在保存/加载期间提供了额外的自定义功能。
1562
+
1563
+ 用于可以被检查点和恢复的对象的有状态协议。
1564
+
1565
+ 从提供的 `state_dict` 恢复对象的状态。
1566
+
1567
+ state_dict (dict[str, Any]) – 要从中恢复的 state dict。
1568
+
1569
+ 对象应将其 `state_dict` 表示形式作为字典返回。此函数的输出将被检查点保存,并稍后在 `load_state_dict()` 中恢复。
1570
+
1571
+ 由于恢复检查点的原位特性,此函数也会在 `torch.distributed.checkpoint.load` 期间被调用。
1572
+
1573
+ 对象状态字典。
1574
+
1575
+ 此示例演示如何使用 Pytorch 分布式检查点保存 FSDP 模型。
1576
+
1577
+ 以下类型定义了检查点期间使用的 IO 接口:
1578
+
1579
+ `load_state_dict` 用来从存储中读取的接口。
1580
+
1581
+ 一个 `StorageReader` 实例在分布式检查点中同时充当协调者和跟随者。作为初始化的一部分,每个实例都会被告知其角色。
1582
+
1583
+ 子类应预期 `load_state_dict` 会进行以下调用序列:
1584
+
1585
+ (所有 rank)如果用户传递了有效的 `checkpoint_id`,则设置 `checkpoint_id`。
1586
+
1587
+ (所有 rank)`read_metadata()`
1588
+
1589
+ (所有 rank)`set_up_storage_reader()`
1590
+
1591
+ (所有 rank)`prepare_local_plan()`
1592
+
1593
+ (协调者)`prepare_global_plan()`
1594
+
1595
+ (所有 rank)`read_data()`
1596
+
1597
+ 执行存储加载的集中规划。
1598
+
1599
+ 此方法仅在协调者实例上调用。
1600
+
1601
+ 虽然此方法可以产生一个完全不同的计划,但首选方法是将存储特定数据存储在 `LoadPlan::storage_data` 中。
1602
+
1603
+ plans (list[torch.distributed.checkpoint.planner.LoadPlan]) – 一个 `LoadPlan` 实例列表,每个 rank 一个。
1604
+
1605
+ 存储全局规划后转换后的 `LoadPlan` 列表
1606
+
1607
+ list[torch.distributed.checkpoint.planner.LoadPlan]
1608
+
1609
+ 执行特定于存储的本地规划。
1610
+
1611
+ 虽然此方法可以产生一个完全不同的计划,但推荐的方法是将存储特定数据存储在 `LoadPlan::storage_data` 中。
1612
+
1613
+ plan (LoadPlan) – 使用中的 `LoadPlan` 的本地计划。
1614
+
1615
+ 存储本地规划后转换后的 `LoadPlan`
1616
+
1617
+ 使用规划器解析数据,并从计划中读取所有项目。
1618
+
1619
+ 子类应调用 `LoadPlanner::load_bytes` 将 `BytesIO` 对象反序列化到正确的位置。
1620
+
1621
+ 子类应调用 `LoadPlanner::resolve_tensor` 以获取应该将数据加载到的张量的访问权限。
1622
+
1623
+ `StorageLayer` 有责任正确调度任何所需的跨设备复制。
1624
+
1625
+ plan (LoadPlan) – 要执行的本地计划
1626
+
1627
+ planner (LoadPlanner) – 用于解析项目的规划器对象。
1628
+
1629
+ 一个在所有读取完成时结束的 Future。
1630
+
1631
+ 读取检查点元数据。
1632
+
1633
+ 与正在加载的检查点关联的元数据对象。
1634
+
1635
+ 调用表示即将进行一次全新的检查点读取。如果用户为此检查点读取设置了 `checkpoint_id`,则可能会存在该 ID。`checkpoint_id` 的含义取决于存储方式。它可以是文件夹/文件的路径,也可以是键值存储的键。
1636
+
1637
+ checkpoint_id (Union[str, os.PathLike, None]) – 此检查点实例的 ID。`checkpoint_id` 的含义取决于存储方式。它可以是文件夹或文件的路径。如果存储更像是一个键值存储,它也可以是一个键。(默认: None)
1638
+
1639
+ 初始化此实例。
1640
+
1641
+ metadata (Metadata) – 要使用的元数据架构。
1642
+
1643
+ is_coordinator (bool) – 此实例是否负责协调检查点。
1644
+
1645
+ 检查给定的 `checkpoint_id` 是否受存储支持。这允许我们启用自动存储选择。
1646
+
1647
+ `save_state_dict` 用来写入存储的接口。
1648
+
1649
+ 一个 `StorageWriter` 实例在分布式检查点中同时充当协调者和跟随者。作为初始化的一部分,每个实例都会被告知其角色。
1650
+
1651
+ 子类应预期会进行以下调用序列。
1652
+
1653
+ (所有 rank)如果用户传递了有效的 `checkpoint_id`,则设置 `checkpoint_id`。
1654
+
1655
+ (所有 rank)`set_up_storage_writer()`
1656
+
1657
+ (所有 rank)`prepare_local_plan()`
1658
+
1659
+ (协调者)`prepare_global_plan()`
1660
+
1661
+ (所有 rank)`write_data()`
1662
+
1663
+ (协调者)`finish()`
1664
+
1665
+ 写入元数据并标记当前检查点为成功。
1666
+
1667
+ 用于序列化元数据的实际格式/架构是一个实现细节。唯一的要求是它可以恢复为相同的对象图。
1668
+
1669
+ metadata (Metadata) – 新检查点的元数据
1670
+
1671
+ results (list[list[torch.distributed.checkpoint.storage.WriteResult]]) – 来自所有 rank 的 `WriteResult` 列表。
1672
+
1673
+ 执行存储的集中规划。
1674
+
1675
+ 此方法仅在协调者实例上调用。
1676
+
1677
+ 虽然此方法可以产生一个完全不同的计划,但首选方法是将存储特定数据存储在 `SavePlan::storage_data` 中。
1678
+
1679
+ plans (list[torch.distributed.checkpoint.planner.SavePlan]) – 一个 `SavePlan` 实例列表,每个 rank 一个。
1680
+
1681
+ 存储全局规划后转换后的 `SavePlan` 列表
1682
+
1683
+ list[torch.distributed.checkpoint.planner.SavePlan]
1684
+
1685
+ 执行特定于存储的本地规划。
1686
+
1687
+ 虽然此方法可以产生一个完全不同的计划,但推荐的方法是将存储特定数据存储在 `SavePlan::storage_data` 中。
1688
+
1689
+ plan (SavePlan) – 使用中的 `SavePlanner` 的本地计划。
1690
+
1691
+ 调用表示即将进行一次全新的检查点写入。如果用户为此检查点写入设置了 `checkpoint_id`,则可能会存在该 ID。`checkpoint_id` 的含义取决于存储方式。它可以是文件夹/文件的路径,也可以是键值存储的键。
1692
+
1693
+ checkpoint_id (Union[str, os.PathLike, None]) – 此检查点实例的 ID。`checkpoint_id` 的含义取决于存储方式。它可以是文件夹或文件的路径。如果存储是键值存储,它也可以是一个键。(默认: None)
1694
+
1695
+ 初始化此实例。
1696
+
1697
+ is_coordinator (bool) – 此实例是否负责协调检查点。
1698
+
1699
+ 返回特定于存储的元数据。这用于在检查点中存储附加信息,这些信息对于提供请求级别的可观察性很有用。`StorageMeta` 在保存调用期间传递给 `SavePlanner`。默认返回 None。
1700
+
1701
+ 示例:
1702
+
1703
+ ```python
1704
+ from torch.distributed.checkpoint.storage import StorageMeta
1705
+
1706
+ class CustomStorageBackend:
1707
+ def get_storage_metadata(self):
1708
+ # 返回将与检查点一起存储的特定于存储的元数据
1709
+ return StorageMeta()
1710
+ ```
1711
+
1712
+ 此示例演示了存储后端如何返回 `StorageMeta`
1713
+ 以将附加元数据附加到检查点。
1714
+
1715
+ 可选[StorageMeta]
1716
+
1717
+ 检查给定的 `checkpoint_id` 是否受存储支持。这允许我们启用自动存储选择。
1718
+
1719
+ 使用规划器解析数据,并从计划中写入所有项目。
1720
+
1721
+ 子类应针对计划中的每个项目调用 `SavePlanner::resolve_data` 以获取要写入的底层对象的访问权限。
1722
+
1723
+ 子类应延迟调用 `resolve_data`,因为它可能会分配内存。对于张量,请做出以下假设:
1724
+
1725
+ 它们可能在任何设备上,包括与 `WriteItem::tensor_data` 上的设备不匹配的情况
1726
+
1727
+ 它们可能是视图或非连续的。只有投影部分需要被保存。
1728
+
1729
+ plan (SavePlan) – 要执行的保存计划。
1730
+
1731
+ planner (SavePlanner) – 用于将项目解析为数据的规划器对象。
1732
+
1733
+ 一个在结束时产生 `WriteResult` 列表的 Future
1734
+
1735
+ Future[list[torch.distributed.checkpoint.storage.WriteResult]]
1736
+
1737
+ 以下类型定义了检查点期间使用的规划器接口:
1738
+
1739
+ 定义 `load_state_dict` 用来规划加载过程的协议的抽象类。
1740
+
1741
+ `LoadPlanner` 是有状态的对象,可用于自定义整个加载过程。
1742
+
1743
+ `LoadPlanner` 充当 `state_dict` 的访问代理,因此对它进行的任何转换将对整个过程可见。
1744
+
1745
+ 在 `load_state_dict` 期间,规划器子类可以预期以下调用序列:
1746
+
1747
+ 发出加载检查点开始的信号。
1748
+
1749
+ 处理 `state_dict` 并生成将被发送以进行全局规划的 `LoadPlan`。
1750
+
1751
+ 获取来自所有 rank 的 `LoadPlan` 并做出任何全局决定。
1752
+
1753
+ 这对于 `state_dict` 中的每个非张量值调用一次。
1754
+
1755
+ 它们对于 `state_dict` 中的每个张量值成对调用。
1756
+
1757
+ 建议用户扩展 `DefaultLoadPlanner` 而不是直接扩展此接口,因为大多数更改都可以通过单个方法的更改来表达。
1758
+
1759
+ 通常有两种扩展模式:
1760
+
1761
+ 重写 `state_dict`。这是扩展加载过程的最简单方法,因为它不需要了解 `LoadPlan` 是如何工作的复杂性。我们需要保留对原始 `state_dict` 的引用,因为加载是原位发生的,所以我们必须能够原位执行它
1762
+
1763
+ 修改 `resolve_tensor` 和 `commit_tensor` 以处理加载时转换。
1764
+
1765
+ 在 `StorageReader` 完成将数据加载到张量中后调用一次。
1766
+
1767
+ 提供的张量与调用 `resolve_tensor` 返回的张量相同。仅当此 `LoadPlanner` 需要在将其复制回 `state_dict` 中的张量之前对张量进行后处理时,才需要此方法。
1768
+
1769
+ 张量的内容将遵循其设备同步模型。
1770
+
1771
+ 计算全局加载计划并返回每个 rank 的计划。
1772
+
1773
+ . 注意:这仅在协调者 rank 上调用
1774
+
1775
+ list[torch.distributed.checkpoint.planner.LoadPlan]
1776
+
1777
+ 根据 `state_dict` 和 `set_up_planner` 提供的元数据创建 `LoadPlan`。
1778
+
1779
+ . 注意:这在每个 rank 上都会被调用。
1780
+
1781
+ 接受来自协调者的计划并返回最终的 `LoadPlan`。
1782
+
1783
+ 加载由 `read_item`` 和 ``value` 描述的项目。
1784
+
1785
+ 此方法预期会原位修改底层的 `state_dict`。
1786
+
1787
+ `value` 的内容由用于生成正在加载的检查点的 `SavePlanner` 定义。
1788
+
1789
+ 返回供 `StorageReader` 加载 `read_item` 时使用的 `BytesIO`。
1790
+
1791
+ `BytesIO` 应与底层 `state_dict` 中的某个对象别名相同,因为 `StorageReader` 将替换其内容。
1792
+
1793
+ 返回供 `StorageReader` 加载 `read_item` 时使用的由 `read_item` 描述的张量。
1794
+
1795
+ 张量应与底层 `state_dict` 中的某个对象别名相同,因为 `StorageReader` 将替换其内容。如果出于任何原因这不可能,规划器可以使用 `commit_tensor` 方法将数据复制回 `state_dict` 中的对象。
1796
+
1797
+ 初始化此实例以将数据加载到 `state_dict` 中。
1798
+
1799
+ . 注意:这在每个 rank 上都会被调用。
1800
+
1801
+ 定义 `save_state_dict` 用来规划保存过程的协议的抽象类。
1802
+
1803
+ `SavePlanner` 是有状态的对象,可用于自定义整个保存过程。
1804
+
1805
+ `SavePlanner` 充当 `state_dict` 的访问代理,因此对它进行的任何转换将对整个过程可见。
1806
+
1807
+ 在 `save_state_dict` 期间,规划器子类可以预期以下调用序列:
1808
+
1809
+ 发出检查点保存开始的信号。
1810
+
1811
+ 处理 `state_dict` 并生成将被发送以进行全局规划的 `SavePlan`。
1812
+
1813
+ 获取来自所有 rank 的 `SavePlan` 并做出任何全局决定。
1814
+
1815
+ 这使得每个 rank 都有机会根据全局规划决策进行调整。
1816
+
1817
+ 在 `state_dict` 中查找供存储层写入的值。
1818
+
1819
+ 建议用户扩展 `DefaultSavePlanner` 而不是直接扩展此接口,因为大多数更改都可以通过单个方法的更改来表达。
1820
+
1821
+ 通常有 3 种扩展模式:
1822
+
1823
+ 重写 `state_dict`。这是扩展保存过程的最简单方法,因为它不需要了解 `SavePlan` 是如何工作的复杂性:
1824
+
1825
+ 串联修改本地计划和查找。这对于需要精细控制数据如何持久化非常有用
1826
+
1827
+ 利用全局规划步骤做出每个 rank 无法单独做出的集中决策
1828
+
1829
+ 最后,一些规划器需要在检查点中保存附加的元数据,这是通过让每个 rank 在本地计划中贡献其数据项,然后由全局规划器聚合它们来实现的:
1830
+
1831
+ 计算全局检查点计划并返回每个 rank 的本地计划。
1832
+
1833
+ 这仅在协调者 rank 上调用。
1834
+
1835
+ tuple[list[torch.distributed.checkpoint.planner.SavePlan], torch.distributed.checkpoint.metadata.Metadata]
1836
+
1837
+ 计算当前 rank 的保存计划。
1838
+
1839
+ 这将被聚合并传递给 `create_global_plan`。特定于规划器的数据可以通过 `SavePlan::planner_data` 传递。
1840
+
1841
+ 这在所有 rank 上调用。
1842
+
1843
+ 合并 `create_local_plan` 创建的计划和 `create_global_plan` 的结果。
1844
+
1845
+ 这在所有 rank 上调用。
1846
+
1847
+ 从 `state_dict` 转换并准备 `write_item` 以供存储,确保幂等性和线程安全。
1848
+
1849
+ 在 `state_dict` 中查找与 `write_item` 关联的对象,并在存储层使用它之前应用任何转换(例如序列化)。
1850
+
1851
+ 在每个 rank 上被多次调用,最终 `SavePlan` 中的每个 `WriteItem` 至少调用一次。
1852
+
1853
+ 此方法应是幂等的且线程安全的。`StorageWriter` 实现可以根据需要自由调用它。
1854
+
1855
+ 任何分配内存的转换都应在此方法被调用时延迟执行,以减少检查点所需的峰值内存。
1856
+
1857
+ 返回张量时,它们可以在任何设备或格式上,也可以是视图。存储层有责任弄清楚如何保存它们。
1858
+
1859
+ Union[Tensor, BytesIO]
1860
+
1861
+ 初始化此规划器以保存 `state_dict`。
1862
+
1863
+ 实现应保存这些值,因为在保存过程的后期将不会提供它们。
1864
+
1865
+ 这在所有 rank 上调用。
1866
+
1867
+ 包含有关需要写入存储的内容信息的数据类。
1868
+
1869
+ 计算底层张量的存储大小,如果这不是张量写入,则为 None。
1870
+
1871
+ 可选[int] 底层张量的存储大小(以字节为单位,如果有)。
1872
+
1873
+ 我们提供了一个基于文件系统的存储层:
1874
+
1875
+ 返回将用于加载检查点的 `checkpoint_id`。
1876
+
1877
+ 使用文件 IO 的 `StorageWriter` 的基本实现。
1878
+
1879
+ 此实现做出以下假设和简化:
1880
+
1881
+ 检查点路径是一个空目录或不存在的目录。
1882
+
1883
+ 文件创建是原子的
1884
+
1885
+ 如果启用了 rank 协调,检查点由每个写入请求一个文件加上一个包含序列化元数据的全局 `.metadata` 文件组成。如果未启用 rank 协调,则由一个 rank 本地的 `__{rank}.metadata` 文件包含序列化元数据。
1886
+
1887
+ 重写 `AsyncStager.stage`
1888
+
1889
+ dict[str, Union[~StatefulT, Any]]
1890
+
1891
+ 我们还提供了其他存储层,包括与 HuggingFace safetensors 交互的存储层:
1892
+
1893
+ .. autoclass:: torch.distributed.checkpoint.HuggingFaceStorageReader :members:
1894
+
1895
+ .. autoclass:: torch.distributed.checkpoint.HuggingFaceStorageWriter :members:
1896
+
1897
+ .. autoclass:: torch.distributed.checkpoint.QuantizedHuggingFaceStorageReader :members:
1898
+
1899
+ 我们提供了 `LoadPlanner` 和 `SavePlanner` 的默认实现,它们可以处理所有 `torch.distributed` 构造,例如 FSDP、DDP、`ShardedTensor` 和 `DistributedTensor`。
1900
+
1901
+ 从规划器接口扩展,以便轻松扩展默认规划器。
1902
+
1903
+ 从规划器接口扩展,以便轻松扩展默认规划器。
1904
+
1905
+ `DefaultLoadPlanner` 在 `LoadPlanner` 之上添加了多项功能。
1906
+
1907
+ 特别是它添加了以下内容:
1908
+
1909
+ flatten_state_dict:处理带有嵌套字典的 state_dict flatten_sharded_tensors:对于二维并行模式下的 FSDP allow_partial_load:如果为 False,当某个键存在于 `state_dict` 中但不在检查点中时,将引发运行时错误。
1910
+
1911
+ 从规划器接口扩展,以便轻松扩展默认规划器。
1912
+
1913
+ 从规划器接口扩展,以便轻松扩展默认规划器。
1914
+
1915
+ 由于早期设计决策,即使原始的未并行化模型完全相同,FSDP 和 DDP 的状态字典可能具有不同的键或全限定名称(例如 `layer1.weight`)。此外,FSDP 提供各种类型的模型状态字典,例如完整和分片状态字典。此外,优化器状态字典使用参数 ID 而不是全限定名称来标识参数,这在使用并行化(例如流水线并行)时可能会导致问题。
1916
+
1917
+ 为了应对这些挑战,我们提供了一套 API,让用户能够轻松管理 state_dict。`get_model_state_dict()` 返回一个模型状态字典,其键与未并行化模型状态字典返回的键一致。类似地,`get_optimizer_state_dict()` 提供的应用于所有并行化中一致的键的优化器状态字典。为了实现这种一致性,`get_optimizer_state_dict()` 会将参数 ID 转换为与未并行化模型状态字典中完全相同的全限定名称。
1918
+
1919
+ 请注意,这些 API 返回的结果可以直接与 `torch.distributed.checkpoint.save()` 和 `torch.distributed.checkpoint.load()` 方法一起使用,无需任何额外转换。
1920
+
1921
+ 提供 `set_model_state_dict()` 和 `set_optimizer_state_dict()` 用于加载由各自获取 API 生成的模型和优化器 state_dict。
1922
+
1923
+ 请注意,只能在优化器上调用 `backward()` 之前或调用 `step()` 之后调用 `set_optimizer_state_dict()`。
1924
+
1925
+ 请注意,此功能是实验性的,API 签名将来可能会发生变化。
1926
+
1927
+ 返回模型 state_dict 和优化器 state_dict。
1928
+
1929
+ `get_state_dict` 可以处理任何由 PyTorch FSDP/fully_shard、DDP/replicate、tensor_parallel/parallelize_module 以及这些并行化的任意组合进行并行化的模块。`get_state_dict` 的主要功能是:1.) 返回可以使用不同数量的训练器和/或不同的并行化进行重新分片的模型和优化器 state_dict。2.) 隐藏特定于并行化的 state_dict API。用户不必调用这些 API。3.) 对结果 state_dict 进行健全性检查。
1930
+
1931
+ 结果状态字典的键是标准的 FQN(全限定名称)。标准 FQN 指的是基于参数在 `nn.Module` 层次结构中位置的 FQN。更具体地说,当模块未被任何并行化分布时,参数的标准 FQN 是 `module.named_parameters()` 或 `module.named_buffers()` 返回的 FQN。由于优化器内部使用参数 ID 来表示参数,因此在调用此 API 时会有一个从参数 ID 到标准 FQN 的转换。
1932
+
1933
+ `get_state_dict` 也可以处理未并行化的模块。在这种情况下,`get_state_dict` 仅执行一项功能 —— 将优化器参数 ID 转换为标准 FQN。
1934
+
1935
+ model (nn.Module) – 模型的 nn.Module。
1936
+
1937
+ optimizers (Union[None, Optimizer, Iterable[Optimizer]]) – 用于优化 `model` 的优化器。
1938
+
1939
+ submodules (已弃用) – 可选[set[nn.Module]]: 仅返回属于这些子模块的模型参数。
1940
+
1941
+ options (StateDictOptions) – 控制如何返回模型 state_dict 和优化器 state_dict 的选项。详情请参见 `StateDictOptions`。
1942
+
1943
+ 包含模型 state_dict 和优化器 state_dict 的元组。
1944
+
1945
+ Tuple[Dict[str, ValueType], OptimizerStateType]
1946
+
1947
+ 返回模型的模型 state_dict。
1948
+
1949
+ 详情请参见 `get_state_dict` 的用法。
1950
+
1951
+ model (nn.Module) – 模型的 nn.Module。
1952
+
1953
+ submodules (已弃用) – 可选[set[nn.Module]]: 仅返回属于这些子模块的模型参数。
1954
+
1955
+ options (StateDictOptions) – 控制如何返回模型 state_dict 和优化器 state_dict 的选项。详情请参见 `StateDictOptions`。
1956
+
1957
+ 模型的 state_dict。
1958
+
1959
+ 返回优化器的组合 state_dict。
1960
+
1961
+ 详情请参见 `get_state_dict` 的用法。
1962
+
1963
+ model (nn.Module) – 模型的 nn.Module。
1964
+
1965
+ optimizers (Union[None, Optimizer, Iterable[Optimizer]]) – 用于优化 `model` 的优化器。
1966
+
1967
+ submodules (已弃用) – 可选[set[nn.Module]]: 仅返回属于这些子模块的模型参数。
1968
+
1969
+ options (StateDictOptions) – 控制如何返回模型 state_dict 和优化器 state_dict 的选项。详情请参见 `StateDictOptions`。
1970
+
1971
+ 优化器的 state_dict。
1972
+
1973
+ 加载模型 state_dict 和优化器 state_dict。
1974
+
1975
+ `get_state_dict` 的对应方法,用于将 state_dict 设置到模型和优化器中。给定的 `model_state_dict` 和 `optim_state_dict` 不必由 `get_state_dict` 返回,但必须满足以下要求:1) 所有 FQN 都是 `get_state_dict` 中定义的标准 FQN,2) 如果张量被分片,它必须是 `ShardedTensor` 或 `DTensor`,3) 优化器 state_dict 不能包含参数 ID;键应该是标准 FQN。
1976
+
1977
+ 在优化器上被调用。否则,优化器状态将无法正确初始化。
1978
+
1979
+ model (nn.Module) – 模型的 nn.Module。
1980
+
1981
+ optimizers (Union[Optimizer, Iterable[Optimizer]]) – 用于优化 `model` 的优化器。
1982
+
1983
+ model_state_dict (Dict[str, ValueType]) – (Union[Dict[nn.Module, Dict[str, ValueType]], Dict[str, ValueType]]): 要加载的模型 state_dict。如果 `model_state_dict` 的键是 `nn.Module`,则该键是 `model` 的子模块,值应为该子模块的 state_dict。加载 state_dict 时,子模块的前缀将被附加到 state_dict 中。
1984
+
1985
+ optim_state_dict (OptimizerStateType) – OptimizerStateType: 要加载的优化器 state_dict。
1986
+
1987
+ options (StateDictOptions) – 控制如何加载模型 state_dict 和优化器 state_dict 的选项。详情请参见 `StateDictOptions`。
1988
+
1989
+ `missing_keys` 是一个字符串列表,包含模型 state_dict 中缺失的键。`unexpected_keys` 是一个字符串列表,包含模型 state_dict 中意外的键。
1990
+
1991
+ `missing_keys` 是一个字符串列表,包含模型 state_dict 中缺失的键。
1992
+
1993
+ `unexpected_keys` 是一个字符串列表,包含模型 state_dict 中意外的键。
1994
+
1995
+ 包含 `missing_keys` 和 `unexpected_keys` 字段的 NamedTuple
1996
+
1997
+ 加载模型 state_dict。
1998
+
1999
+ `get_model_state_dict` 的对应方法,用于将 state_dict 设置到模型中。详情请参见 `set_state_dict` 的用法。
2000
+
2001
+ model (nn.Module) – 模型的 nn.Module。
2002
+
2003
+ model_state_dict (Dict[str, ValueType]) – (Dict[str, ValueType]): 要加载的模型 state_dict。如果 `model_state_dict` 的键是 `nn.Module`,则该键是 `model` 的子模块,值应为该子模块的 state_dict。加载 state_dict 时,子模块的前缀将被附加到 state_dict 中。
2004
+
2005
+ options (StateDictOptions) – 控制如何加载模型 state_dict 和优化器 state_dict 的选项。详情请参见 `StateDictOptions`。
2006
+
2007
+ `missing_keys` 是一个字符串列表,包含缺失的键 `unexpected_keys` 是一个字符串列表,包含意外的键
2008
+
2009
+ `missing_keys` 是一个字符串列表,包含缺失的键
2010
+
2011
+ `unexpected_keys` 是一个字符串列表,包含意外的键
2012
+
2013
+ 包含 `missing_keys` 和 `unexpected_keys` 字段的 NamedTuple
2014
+
2015
+ 加载优化器 state_dict。
2016
+
2017
+ `get_optimizer_state_dict` 的对应方法,用于将 state_dict 设置到优化器中。详情请参见 `set_state_dict` 的用法。
2018
+
2019
+ 在优化器上调用 `step()`。否则,优化器状态将无法正确初始化。
2020
+
2021
+ model (nn.Module) – 模型的 nn.Module。
2022
+
2023
+ optimizers (Union[Optimizer, Iterable[Optimizer]]) – 用于优化 `model` 的优化器。
2024
+
2025
+ optim_state_dict (OptimizerStateType) – OptimizerStateType: 要加载的优化器 state_dict。
2026
+
2027
+ options (StateDictOptions) – 控制如何加载模型 state_dict 和优化器 state_dict 的选项。详情请参见 `StateDictOptions`。
2028
+
2029
+ 此数据类指定 `get_state_dict`/`set_state_dict` 将如何工作。
2030
+
2031
+ full_state_dict: 如果设置为 True,则返回的 state_dict 中的所有张量将被收集。返回的 state_dict 中将不存在 ShardedTensor 和 DTensor。
2032
+
2033
+ cpu_offload: 将所有张量卸载到 CPU。为了防止 CPU 发生 OOM,如果 `full_state_dict` 也为 true,则只有 rank0 将获取 state_dict,而所有其他 rank 将获取空的 state_dict。
2034
+
2035
+ ignore_frozen_params: 如果值为 True,则返回的 state_dict 将不包含任何冻结的参数 —— 即 `requires_grad` 为 False 的参数。默认值为 False。
2036
+
2037
+ keep_submodule_prefixes (已弃用): 当 `submodules` 不为 None 时,此选项指示是否保留 state_dict 键中的子模块前缀。例如,如果子模块是 `module.pretrain` 并且参数的完整 FQN 是参数的 `pretrain.layer1.weight`。当此选项为 True 时,返回的 state_dict 中参数的键将是 `pretrain.layer1.weight`。如果选项为 False,则键将是 `layer1.weight`。请注意,如果 `keep_submodule_prefixes` 为 False,可能会出现冲突的 FQN,因此 `submodules` 中应该只有一个子模块。
2038
+
2039
+ strict: `set_state_dict` 调用 `model.load_state_dict()` 时的严格选项。
2040
+
2041
+ 完整的 state_dict,并会将 state_dict/optim_state_dict 中的张量逐一广播给其他 rank。其他 rank 将接收张量并根据模型和优化器中的本地分片进行分片。使用此选项时必须将 `full_state_dict` 设置为 True。此选项目前仅支持 DTensor,不支持遗留的 ShardedTensor。
2042
+
2043
+ 对于习惯使用和共享 `torch.save` 格式的模型的用户,提供了以下方法,它们提供用于在格式之间进行转换的离线实用程序。
2044
+
2045
+ 给定一个
2046
+
2047
+ ## torch.distributed.tensor#
2048
+
2049
+ **URL:** https://pytorch.org/docs/stable/distributed.tensor.html
2050
+
2051
+ **目录:**
2052
+ - torch.distributed.tensor#
2053
+ - PyTorch DTensor (Distributed Tensor)#
2054
+ - DTensor Class APIs#
2055
+ - DeviceMesh as the distributed communicator#
2056
+ - DTensor Placement Types#
2057
+ - Different ways to create a DTensor#
2058
+ - Create DTensor from a logical torch.Tensor#
2059
+ - DTensor Factory Functions#
2060
+ - Random Operations#
2061
+ - Debugging#
2062
+
2063
+ 创建时间:2025 年 6 月 13 日 | 最后更新时间:2025 年 8 月 23 日
2064
+
2065
+ torch.distributed.tensor 目前处于 alpha 状态并在开发中,我们承诺对文档中列出的大多数 API 保持向后兼容,但如有必要,可能会发生 API 变更。
2066
+
2067
+ PyTorch DTensor 提供了简单灵活的张量切分原语,可透明地处理分布式逻辑,包括分片存储、算子计算以及跨设备/主机的集合通信。DTensor 可用于构建不同的并行解决方案,并支持在处理多维分片时的分片 state_dict 表示。
2068
+
2069
+ 请查看基于 DTensor 构建的 PyTorch 原生并行解决方案的示例:
2070
+
2071
+ DTensor 遵循 SPMD(单程序多数据)编程模型,使用户能够像编写单设备程序一样编写分布式程序,并具有相同的收敛属性。它通过指定 DeviceMesh 和 Placement 来提供统一的张量分片布局(DTensor Layout):
2072
+
2073
+ DeviceMesh 使用 n 维数组表示集群的设备拓扑和通信器。
2074
+
2075
+ Placement 描述逻辑张量在 DeviceMesh 上的分片布局。DTensor 支持三种类型的放置方式:Shard(分片)、Replicate(复制)和 Partial(部分)。
2076
+
2077
+ DTensor 是 torch.Tensor 的子类。这意味着一旦创建了 DTensor,就可以像使用 torch.Tensor 一样使用它,包括像在单设备上运行一样运行不同类型的 PyTorch 算子,从而允许 PyTorch 算子进行正确的分布式计算。
2078
+
2079
+ 除了现有的 torch.Tensor 方法外,它还提供了一组附加方法来与 torch.Tensor 交互、将 DTensor Layout 重新分配为新的 DTensor、在所有设备上获取完整的张量内容等。
2080
+
2081
+ DTensor (Distributed Tensor) 是 torch.Tensor 的子类,它提供类似单设备的抽象,以便对多设备 torch.Tensor 进行编程。它通过 DeviceMesh 和以下类型的 Placement 来描述分布式张量分片布局(DTensor Layout):
2082
+
2083
+ Shard:在 DeviceMesh 维度对应的设备上,对张量的维度 dim 进行切分
2084
+
2085
+ Replicate:在 DeviceMesh 维度的设备上复制张量
2086
+
2087
+ Partial:在 DeviceMesh 维度的设备上,张量处于待归约状态
2088
+
2089
+ 在调用 PyTorch 算子时,DTensor 会重写 PyTorch 算子以执行分片计算,并在必要时发起通信。伴随着算子计算,DTensor 会(根据算子本身的语义)正确地转换或传播放置方式(DTensor Layout),并生成新的 DTensor 输出。
2090
+
2091
+ 为了确保在调用 PyTorch 算子时 DTensor 分片计算的数值正确性,DTensor 要求该算子的每个张量参数都必须是 DTensor。
2092
+
2093
+ 在此直接使用 Tensor 子类构造函数并不是创建 DTensor 的推荐方法(即它不能正确处理自动求导,因此不是公开的 API)。请参阅 create_dtensor 部分以了解如何创建 DTensor。
2094
+
2095
+ 返回一个 ChunkStorageMetadata 列表,它是一个描述当前 rank 上本地分片/副本的大小/偏移量的数据类。对于 DTensor,每个 rank 将有一个本地分片/副本,因此返回的列表通常只有一个元素。
2096
+
2097
+ 此 dunder 方法主要用于分布式检查点目的。
2098
+
2099
+ 一个表示当前 rank 上分片大小/偏移量的 List[ChunkStorageMetadata] 对象。
2100
+
2101
+ 根据指定的 device_mesh 和 placements,在每个 rank 上通过本地 torch.Tensor 创建一个 DTensor。
2102
+
2103
+ local_tensor (torch.Tensor) – 每个 rank 上的本地 torch.Tensor。
2104
+
2105
+ device_mesh (DeviceMesh, 可选) – 用于放置张量的 DeviceMesh,如果未指定,必须在 DeviceMesh 上下文管理器下调用,默认值:None
2106
+
2107
+ placements (List[Placement], 可选) – 描述如何将本地 torch.Tensor 放置到 DeviceMesh 上的放置方式,其元素数量必须与 device_mesh.ndim 相同。
2108
+
2109
+ run_check (bool, 可选) – 以额外的通信为代价,跨 rank 执行完整性检查,检查每个本地张量的元信息以确保正确性。如果 placements 中包含 Replicate,设备网格维度上第一个 rank 的数据将被广播到其他 rank。默认值:False
2110
+
2111
+ shape (torch.Size, 可选) – 一个整数列表,指定建立在 local_tensor 之上的 DTensor 的大小。注意,如果各个 rank 上的 local_tensor 形状不同,则需要提供此参数。如果未提供,将假设给定的分布式张量在各个 rank 上被均匀切分来计算形状。默认值:None
2112
+
2113
+ stride (tuple, 可选) – 一个整数列表,指定 DTensor 的步长。如果未提供,将假设给定的分布式张量在各个 rank 上被均匀切分来计算步长。默认值:None
2114
+
2115
+ 当 run_check=False 时,用户有责任确保跨 rank 传入的本地张量是正确的(即对于 Shard(dim) 放置方式,张量应被分片;对于 Replicate() 放置方式,张量应被复制)。否则,创建的 DTensor 的行为是未定义的。
2116
+
2117
+ from_local 是可微的,创建的 DTensor 对象的 requires_grad 将取决于 local_tensor 是否需要梯度。
2118
+
2119
+ 返回此 DTensor 的完整张量。它将执行必要的集合通信,从其 DeviceMesh 上的其他 rank 收集本地张量并将它们拼接在一起。它是以下代码的语法糖:
2120
+
2121
+ dtensor.redistribute(placements=[Replicate()] * mesh.ndim).to_local()
2122
+
2123
+ grad_placements (List[Placement], 可选) – 该放置方式描述了此函数返回的完整张量的任何梯度布局的未来布局。full_tensor 将 DTensor 转换为完整的 torch.Tensor,并且返回的 torch.tensor 稍后在代码中可能不会用作原始复制的 DTensor 布局。此参数是用户可以在返回张量的梯度布局与原始复制 DTensor 布局不匹配时提供给自动求导的提示。如果未指定,我们将假设完整张量的梯度布局为复制。
2124
+
2125
+ 一个表示此 DTensor 完整张量的 torch.Tensor 对象。
2126
+
2127
+ full_tensor 是可微的。
2128
+
2129
+ redistribute 执行必要的集合操作,将当前 DTensor 从其当前的 placements 重新分配为新的 placements,或者从其当前的 DeviceMesh 重新分配到新的 DeviceMesh。即,我们可以通过为 DeviceMesh 的每个维度指定 Replicate 放置方式,将分片的 DTensor 转换为复制的 DTensor。
2130
+
2131
+ 在一个设备网格维度上从当前 placements 重新分配到新的 placements 时,我们将执行以下包含通信集合操作或本地操作:
2132
+
2133
+ Shard(dim) -> Replicate(): all_gather
2134
+
2135
+ Shard(src_dim) -> Shard(dst_dim): all_to_all
2136
+
2137
+ Replicate() -> Shard(dim): 本地分块(即 torch.chunk)
2138
+
2139
+ Partial() -> Replicate(): all_reduce
2140
+
2141
+ Partial() -> Shard(dim): reduce_scatter
2142
+
2143
+ redistribute 能够正确地为在 1 维或 N 维 DeviceMesh 上创建的 DTensor 找出必要的重新分配步骤。
2144
+
2145
+ device_mesh (DeviceMesh, 可选) – 用于放置 DTensor 的 DeviceMesh。如果未指定,它将使用当前 DTensor 的 DeviceMesh。默认值:None
2146
+
2147
+ placements (List[Placement], 可选) – 描述如何将 DTensor 放置到 DeviceMesh 中的新放置方式,其元素数量必须与 device_mesh.ndim 相同。默认值:在所有网格维度上进行复制
2148
+
2149
+ async_op (bool, 可选) – 是否异步执行 DTensor 重新分配操作。默认: False
2150
+
2151
+ forward_dtype (torch.dtype, 可选) – 在其前向传播中重新分配本地张量之前,可以将本地张量数据类型转换为 forward_dtype。生成的 DTensor 将处于 forward_dtype 默认: None。
2152
+
2153
+ backward_dtype (torch.dtype, 可选) – 在其反向传播中重新分配本地张量之前,可以将本地张量数据类型转换为 backward_dtype。生成的 DTensor 梯度将被转换回当前 DTensor 的数据类型。默认: None
2154
+
2155
+ redistribute 是可微的,这意味着用户不需要担心 redistribute 操作的反向传播公式。
2156
+
2157
+ redistribute 目前仅支持在同一个 DeviceMesh 上重新分配 DTensor,如果您需要将 DTensor 重新分配到不同的 DeviceMesh,请提交一个 issue。
2158
+
2159
+ 获取此 DTensor 在其当前 rank 上的本地张量。对于分片,它返回逻辑张量视图的一个本地分片;对于复制,它返回其当前 rank 上的副本。
2160
+
2161
+ grad_placements (List[Placement], 可选) – 该放置方式描述了此函数返回的张量的任何梯度布局的未来布局。to_local 将 DTensor 转换为本地张量,并且返回的本地张量稍后在代码中可能不会用作原始 DTensor 布局。此参数是用户可以在返回张量的梯度布局与原始 DTensor 布局不匹配时提供给自动求导的提示。如果未指定,我们将假设梯度布局保持与原始 DTensor 相同,并将其用于梯度计算。
2162
+
2163
+ 一个 torch.Tensor 或 AsyncCollectiveTensor 对象。它表示其当前 rank 上的本地张量。当返回 AsyncCollectiveTensor 对象时,意味着本地张量尚未就绪(即通信尚未完成)。在这种情况下,用户需要调用 wait 等待本地张量就绪。
2164
+
2165
+ to_local 是可微的,返回的本地张量的 requires_grad 将取决于该 DTensor 是否需要梯度。
2166
+
2167
+ 与此 DTensor 对象关联的 DeviceMesh 属性。
2168
+
2169
+ device_mesh 是只读属性,不能被设置。
2170
+
2171
+ 此 DTensor 的 placements 属性,描述了该 DTensor 在其 DeviceMesh 上的布局。
2172
+
2173
+ placements 是只读属性,不能被设置。
2174
+
2175
+ DeviceMesh 从 DTensor 构建而来,作为描述集群设备拓扑并表示多维通信器(构建在 ProcessGroup 之上)的抽象。要详细了解如何创建/使用 DeviceMesh,请参阅 DeviceMesh 指南。
2176
+
2177
+ DTensor 在每个 DeviceMesh 维度上支持以下类型的 Placement:
2178
+
2179
+ Shard(dim) 放置方式描述了 DTensor 在对应的 DeviceMesh 维度上对张量维度 dim 的切分,其中 DeviceMesh 维度上的每个 rank 仅持有全局张量的一个分片/部分。Shard(dim) 放置方式遵循 torch.chunk(dim) 语义,其中当张量维度无法在 DeviceMesh 维度上被整除时,DeviceMesh 维度上的最后几个分片可能为空。Shard 放置方式可以被所有的 DTensor API 使用(即 distribute_tensor、from_local 等)
2180
+
2181
+ dim (int) – 描述 DTensor 在其对应的 DeviceMesh 维度上进行切分的张量维度。
2182
+
2183
+ 当张量维度的大小无法在 DeviceMesh 维度上整除时,对该张量维度进行分片目前是实验性的,可能会发生变化。
2184
+
2185
+ Replicate() 放置方式描述了 DTensor 在对应的 DeviceMesh 维度上进行复制,其中 DeviceMesh 维度上的每个 rank 都持有全局张量的一个副本。Replicate 放置方式可以被所有的 DTensor API 使用(即 distribute_tensor、DTensor.from_local 等)
2186
+
2187
+ Partial(reduce_op) 放置方式描述了在指定的 DeviceMesh 维度上处于待归约状态的 DTensor,其中 DeviceMesh 维度上的每个 rank 都持有全局张量的部分值。用户可以使用 redistribute 将 Partial DTensor 重新分配为指定 DeviceMesh 维度上的 Replicate 或 Shard(dim) 放置方式,这将在底层触发必要的通信操作(即 allreduce、reduce_scatter)。
2188
+
2189
+ reduce_op (str, 可选) – 用于部分 DTensor 生成复制/分片 DTensor 的归约操作。仅支持逐元素归约操作,包括:"sum"、"avg"、"product"、"max"、"min",默认值:"sum"。
2190
+
2191
+ Partial 放置方式可以作为 DTensor 算子的结果生成,并且只能被 DTensor.from_local API 使用。
2192
+
2193
+ Placement 类型的基类,它描述了 DTensor 是如何放置到 DeviceMesh 上的。Placement 和 DeviceMesh 一起可以描述 DTensor Layout。它是三种主要 DTensor Placement 类型的基类:Shard、Replicate 和 Partial。
2194
+
2195
+ 此类不适合直接使用,主要作为类型提示存根。
2196
+
2197
+ distribute_tensor() 在每个 rank 上通过逻辑或“全局” torch.Tensor 创建一个 DTensor。这可用于对叶子 torch.Tensor(即模型参数/缓冲区和输入)进行切分。
2198
+
2199
+ DTensor.from_local() 在每个 rank 上通过本地 torch.Tensor 创建一个 DTensor,这可用于从非叶子 torch.Tensor(即前向/反向传播过程中的中间激活张量)创建 DTensor。
2200
+
2201
+ DTensor 提供了专用的张量工厂函数(例如 empty()、ones()、randn() 等),允许通过直接指定 DeviceMesh 和 Placement 进行不同的 DTensor 创建。与 distribute_tensor() 相比,这可以直接在设备上实例化分片内存,而不是在初始化逻辑张量内存之后再执行分片。
2202
+
2203
+ torch.distributed 中的 SPMD(单程序多数据)编程模型会启动多个进程(例如通过 torchrun)来执行同一个程序,这意味着程序中的模型将首先在不同的进程上被初始化(即模型可能在 CPU、或 meta 设备上初始化,或者在内存充足的情况下直接在 GPU 上初始化)。
2204
+
2205
+ DTensor 提供了一个 distribute_tensor() API,可以将模型权重或张量切分为多个 DTensor,它会在每个进程上通过“逻辑”张量创建一个 DTensor。这将使创建的 DTensor 遵循单设备语义,这对于数值正确性至关重要。
2206
+
2207
+ 根据指定的 placements 将叶子 torch.Tensor(即 nn.Parameter/缓冲区)分发到 device_mesh。device_mesh 和 placements 的秩必须相同。要分发的张量是逻辑或“全局”张量,并且该 API 将使用 DeviceMesh 维度上第一个 rank 的张量作为真实数据源,以保留单设备语义。如果您想在自动求导计算过程中构建 DTensor,请改用 DTensor.from_local()。
2208
+
2209
+ tensor (torch.Tensor) – 要分发的 torch.Tensor。请注意,如果您想在一个不能被该网格维度中的设备数整除的维度上对张量进行切分,我们将使用 torch.chunk 语义对张量进行切分并分散这些分片。这种不均匀分片的行为是实验性的,可能会发生变化。
2210
+
2211
+ device_mesh (DeviceMesh, 可选) – 用于分发张量的 DeviceMesh,如果未指定,必须在 DeviceMesh 上下文管理器下调用,默认值:None
2212
+
2213
+ placements (List[Placement], 可选) – 描述如何将张量放置到 DeviceMesh 上的放置方式,其元素数量必须与 device_mesh.ndim 相同。如果未指定,我们默认将跨越 device_mesh 从每个维度的第一个 rank 复制张量。
2214
+
2215
+ src_data_rank (int, 可选) – 逻辑/全局张量的源数据 rank,distribute_tensor() 使用它将分片/副本分散/广播到其他 rank。默认情况下,我们在每个 DeviceMesh 维度上使用 group_rank=0 作为源数据,以保留单设备语义。如果显式传入 None,distribute_tensor() 将仅使用其本地数据,而不是尝试通过分散/广播来保留单设备语义。默认: 0
2216
+
2217
+ 一个 DTensor 或 XLAShardedTensor 对象。
2218
+
2219
+ 当使用 xla 设备类型初始化 DeviceMesh 时,distribute_tensor 返回的是 XLAShardedTensor。有关更多详细信息,请参见此 issue。XLA 集成是实验性的,可能会发生变化。
2220
+
2221
+ 除了 distribute_tensor() 之外,DTensor 还提供了一个 distribute_module() API,以便更容易在 nn.Module 级别上进行分片。
2222
+
2223
+ 该函数公开了三个函数来控制模块的参数/输入/输出:
2224
+
2225
+ 1. 通过指定 partition_fn 在运行时执行之前对模块执行分片(即允许用户根据指定的 partition_fn 将模块参数转换为 DTensor 参数)。 2. 通过指定 input_fn 和 output_fn 控制运行时执行期间模块的输入或输出。(即把输入转换为 DTensor,将输出转换回 torch.Tensor)
2226
+
2227
+ module (nn.Module) – 要被分区的用户模块。
2228
+
2229
+ device_mesh (DeviceMesh) – 用于放置模块的设备网格。
2230
+
2231
+ partition_fn (Callable) – 对参数进行分区的函数(即在 device_mesh 上对某些参数进行分片)。如果未指定 partition_fn,默认情况下我们将在网格上复制该模块的所有模块参数。
2232
+
2233
+ input_fn (Callable) – 指定输入分布,例如可以控制模块的输入如何被分片。input_fn 将作为模块的 forward_pre_hook(前向前置钩子)安装。
2234
+
2235
+ output_fn (Callable) – 指定输出分布,例如可以控制输出如何被分片,或将其转换回 torch.Tensor。output_fn 将作为模块的 forward_hook(前向后置钩子)安装。
2236
+
2237
+ 一个包含全为 DTensor 参数/缓冲区的模块。
2238
+
2239
+ 当使用 xla 设备类型初始化 DeviceMesh 时,distribute_module 返回带有 PyTorch/XLA SPMD 注解参数的 nn.Module。有关更多详细信息,请参见此 issue。XLA 集成是实验性的,可能会发生变化。
2240
+
2241
+ DTensor 还提供了专用的张量工厂函数,允许通过额外指定所创建 DTensor 的 DeviceMesh 和 Placement,直接使用类似 torch.Tensor 的工厂函数 API(即 torch.ones、torch.empty 等)来创建 DTensor:
2242
+
2243
+ 返回一个填充有标量值 0 的 DTensor。
2244
+
2245
+ size (int...) – 定义输出 DTensor 形状的整数序列。可以是可变数量的参数或像列表或元组这样的集合。例如:zeros(1,2,3..) 或 zeros([1,2,3..]) 或 zeros((1,2,3..))
2246
+
2247
+ requires_grad (bool, 可选) – 自动求导是否应记录在返回的 DTensor 上的操作。默认: False。
2248
+
2249
+ dtype (torch.dtype, 可选) – 返回的 DTensor 所需的数据类型。默认: 如果为 None,则使用全局默认值(参见 torch.set_default_dtype())。
2250
+
2251
+ layout (torch.layout, 可选) – 返回的 DTensor 所需的布局。默认: torch.strided。
2252
+
2253
+ device_mesh – DeviceMesh 类型,包含 rank 的网格信息。
2254
+
2255
+ placements – Placement 类型的序列:Shard、Replicate。
2256
+
2257
+ 每个 rank 上的一个 DTensor 对象。
2258
+
2259
+ 返回一个填充有标量值 1 的 DTensor,其形状由可变参数 size 定义。
2260
+
2261
+ size (int...) – 定义输出 DTensor 形状的整数序列。可以是可变数量的参数或像列表或元组这样的集合。例如:ones(1,2,3..) 或 ones([1,2,3..]) 或 ones((1,2,3..))
2262
+
2263
+ dtype (torch.dtype, 可选) – 返回的 DTensor 所需的数据类型。默认: 如果为 None,则使用全局默认值(参见 torch.set_default_dtype())。
2264
+
2265
+ layout (torch.layout, 可选) – 返回的 DTensor 所需的布局。默认: torch.strided。
2266
+
2267
+ requires_grad (bool, 可选) – 自动求导是否应记录在返回的 DTensor 上的操作。默认: False。
2268
+
2269
+ device_mesh – DeviceMesh 类型,包含 rank 的网格信息。
2270
+
2271
+ placements – Placement 类型的序列:Shard、Replicate。
2272
+
2273
+ 每个 rank 上的一个 DTensor 对象。
2274
+
2275
+ 返回一个用未初始化数据填充的 DTensor。该 DTensor 的形状由可变参数 size 定义。
2276
+
2277
+ size (int...) – 定义输出 DTensor 形状的整数序列。可以是可变数量的参数或像列表或元组这样的集合。例如:empty(1,2,3..) 或 empty([1,2,3..]) 或 empty((1,2,3..))
2278
+
2279
+ dtype (torch.dtype, 可选) – 返回的 DTensor 所需的数据类型。默认: 如果为 None,则使用全局默认值(参见 torch.set_default_dtype())。layout (torch.layout, 可选): 返回的 DTensor 所需的布局。默认: torch.strided。
2280
+
2281
+ requires_grad (bool, 可选) – 自动求导是否应记录在返回的 DTensor 上的操作。默认: False。
2282
+
2283
+ device_mesh – DeviceMesh 类型,包含 rank 的网格信息。
2284
+
2285
+ placements – Placement 类型的序列:Shard、Replicate。
2286
+
2287
+ 每个 rank 上的一个 DTensor 对象。
2288
+
2289
+ 根据 device_mesh 和 placements 返回一个用 fill_value 填充的 DTensor,其形状由参数 size 定义。
2290
+
2291
+ size (int...) – 定义输出 DTensor 形状的整数序列。可以是可变数量的参数或像列表或元组这样的集合。例如:ones(1,2,3..) 或 ones([1,2,3..]) 或 ones((1,2,3..))
2292
+
2293
+ fill_value (Scalar) – 用于填充输出张量的值。
2294
+
2295
+ dtype (torch.dtype, 可选) – 返回的 DTensor 所需的数据类型。默认: 如果为 None,则使用全局默认值(参见 torch.set_default_dtype())。
2296
+
2297
+ layout (torch.layout, 可选) – 返回的 DTensor 所需的布局。默认: torch.strided。
2298
+
2299
+ requires_grad (bool, 可选) – 自动求导是否应记录在返回的 DTensor 上的操作。默认: False。
2300
+
2301
+ device_mesh – DeviceMesh 类型,包含 rank 的网格信息。
2302
+
2303
+ placements – Placement 类型的序列:Shard、Replicate。
2304
+
2305
+ 每个 rank 上的一个 DTensor 对象。
2306
+
2307
+ 返回一个用 [0, 1) 区间内均匀分布的随机数填充的 DTensor。张量的形状由可变参数 size 定义。
2308
+
2309
+ size (int...) – 定义输出 DTensor 形状的整数序列。可以是可变数量的参数或像列表或元组这样的集合。例如:ones(1,2,3..) 或 ones([1,2,3..]) 或 ones((1,2,3..))
2310
+
2311
+ dtype (torch.dtype, 可选) – 返回的 DTensor 所需的数据类型。默认: 如果为 None,则使用全局默认值(参见 torch.set_default_dtype())。
2312
+
2313
+ layout (torch.layout, 可选) – 返回的 DTensor 所需的布局。默认: torch.strided。
2314
+
2315
+ requires_grad (bool, 可选) – 自动求导是否应记录在返回的 DTensor 上的操作。默认: False。
2316
+
2317
+ device_mesh – DeviceMesh 类型,包含 rank 的网格信息。
2318
+
2319
+ placements – Placement 类型的序列:Shard、Replicate。
2320
+
2321
+ 每个 rank 上的一个 DTensor 对象。
2322
+
2323
+ 返回一个用均值为 0、方差为 1 的正态分布中提取的随机数填充的 DTensor。张量的形状由可变参数 size 定义。
2324
+
2325
+ size (int...) – 定义输出 DTensor 形状的整数序列。可以是可变数量的参数或像列表或元组这样的集合。例如:ones(1,2,3..) 或 ones([1,2,3..]) 或 ones((1,2,3..))
2326
+
2327
+ dtype (torch.dtype, 可选) – 返回的 DTensor 所需的数据类型。默认: 如果为 None,则使用全局默认值(参见 torch.set_default_dtype())。
2328
+
2329
+ layout (torch.layout, 可选) – 返回的 DTensor 所需的布局。默认: torch.strided。
2330
+
2331
+ requires_grad (bool, 可选) – 自动求导是否应记录在返回的 DTensor 上的操作。默认: False。
2332
+
2333
+ device_mesh – DeviceMesh 类型,包含 rank 的网格信息。
2334
+
2335
+ placements – Placement 类型的序列:Shard、Replicate。
2336
+
2337
+ 每个 rank 上的一个 DTensor 对象。
2338
+
2339
+ DTensor 提供了分布式 RNG 功能,以确保对分片张量执行的随机操作能获得唯一的值,而对复制张量执行的随机操作能获得相同的值。该系统要求所有参与的 rank(例如 SPMD rank)在执行每次 dtensor 随机操作之前,都以相同的生成器状态开始;如果满足此条件,它将确保在每次 dtensor 随机操作完成后,它们最终都达到相同的状态。在随机操作期间不执行任何通信来同步 RNG 状态。
2340
+
2341
+ 接受 generator 关键字参数的算子将使用用户传入的生成器(如果传入了的话),否则使用该设备的默认生成器。无论使用哪种生成器,它都会在 DTensor 操作之后更新。将同一个生成器同时用于 DTensor 和非 DTensor 操作是有效的,但在此情况下,必须注意确保非 DTensor 操作在所有 rank 上同等地更新生成器状态。
2342
+
2343
+ 当将 DTensor 与流水线并行结合使用时,每个流水线阶段的 rank 应使用不同的种子,而流水线阶段内的 rank 应使用相同的种子。
2344
+
2345
+ DTensor 的 RNG 基础设施基于 philox 的 RNG 算法,并支持任何基于 philox 的后端(cuda 及其他类 cuda 设备),但遗憾的是尚不支持 CPU 后端。
2346
+
2347
+ 启动程序时,您可以使用来自 torch._logging 的 TORCH_LOGS 环境变量开启额外的日志记录:
2348
+
2349
+ TORCH_LOGS=+dtensor 将显示 logging.DEBUG 消息及以上的所有级别。
2350
+
2351
+ TORCH_LOGS=dtensor 将显示 logging.INFO 及以上的消息。
2352
+
2353
+ TORCH_LOGS=-dtensor 将显示 logging.WARNING 及以上的消息。
2354
+
2355
+ 为了调试应用了 DTensor 的程序,并了解底层发生了哪些集合通信的更多细节,DTensor 提供了 CommDebugMode:
2356
+
2357
+ CommDebugMode 是一个上下文管理器,用于统计其上下文内的功能集合通信数量。它是通过 TorchDispatchMode 来实现这一点的。
2358
+
2359
+ 目前尚未支持所有的集合通信。
2360
+
2361
+ 生成显示模块级操作和集合通信追踪信息的详细表格。信息量取决于 noise_level。
2362
+
2363
+ 打印模块级别的集合通信计数
2364
+
2365
+ 打印不包含在平凡操作中的 dTensor 操作,以及模块信息
2366
+
2367
+ 打印不包含在平凡操作中的操作
2368
+
2369
+ 打印所有操作
2370
+
2371
+ 创建用于构建浏览器可视化的 json 文件。0. 打印模块级别的集合通信计数 1. 打印不包含在平凡操作中的 dTensor 操作 2. 打印不包含在平凡操作中的操作 3. 打印所有操作。
2372
+
2373
+ 以字典形式返回通信计数。
2374
+
2375
+ 以字典形式表示的通信计数。
2376
+
2377
+ dict[str, dict[str, Any]]
2378
+
2379
+ dict[str, dict[str, Any]]
2380
+
2381
+ 作为控制台 CommDebugMode 输出的替代方案,将内容写入由用户指定的文件中。
2382
+
2383
+ 为了可视化维度少于 3 维的 DTensor 的分片情况,DTensor 提供了 visualize_sharding():
2384
+
2385
+ 在终端中可视化 1 维或 2 维 DTensor 的分片。
2386
+
2387
+ 此操作需要安装 tabulate 包,或者 rich 和 matplotlib 包。对于空张量,不会打印任何分片信息。
2388
+
2389
+ DTensor 还提供了一系列实验性功能。这些功能要么处于原型阶段,要么基本功能已经完成但正在寻找用户反馈。如果您对这些功能有反馈意见,请向 PyTorch 提交 issue。
2390
+
2391
+ context_parallel 是一个用于启用上下文并行 (CP) 的实验性 API。此 API 执行两个操作:1)使用支持 CP 的版本来修补 SDPA (torch.nn.functional.scaled_dot_product_attention),2)沿序列维度对缓冲区进行分片,每个 rank 将根据网格保留相应的分片。
2392
+
2393
+ mesh (DeviceMesh) – 用于上下文并行的设备网格。
2394
+
2395
+ buffers (可选[List[torch.Tensor]]) – 其使用依赖于序列维度的缓冲区。示例包括输入批次、标签和位置嵌入缓冲区。这些缓冲区必须沿序列维度进行分片以确保准确性。分片将就地发生,缓冲区的形状将在上下文内改变。上下文结束后缓冲区将被恢复。no_restore_buffers 可用于指定哪些缓冲区不需要恢复。请注意,buffers 不应包含任何 nn.Parameter。
2396
+
2397
+ buffer_seq_dims (可选[List[int]]) – 缓冲区的序列维度。
2398
+
2399
+ no_restore_buffers (可选[Set[torch.Tensor]]) – 此集合中的缓冲区在上下文退出后将不会被恢复。此集合必须是 buffers 的子集。如果缓冲区在上下文退出后不再使用,可以将这些缓冲区放入此列表中以避免额外的恢复时间。
2400
+
2401
+ Generator[None, None, None]
2402
+
2403
+ torch.distributed.tensor.experimental.context_parallel 是 PyTorch 中的一个原型功能。该 API 可能会发生变化。
2404
+
2405
+ local_map() 是一个实验性 API,允许用户将 DTensor 传递给为应用于 torch.Tensor 而编写的函数。它的实现方式是提取 DTensor 的本地组件,调用该函数,然后根据 out_placements 将输出包装回 DTensor。
2406
+
2407
+ func (Callable) – 应用于每个 DTensor 本地分片的函数。
2408
+
2409
+ out_placements (Union[PlacementType, Tuple[PlacementType, …]]) – func 展平输出中 DTensor 所需的放置方式。如果展平的输出是单个值,则 out_placements 应为 PlacementType 类型。如果展平的输出包含多个值,则 out_placements 应为 PlacementType 值的元组,并与展平的输出 1:1 映射。此外,对于 Tensor 输出,我们使用 PlacementType 作为其放置方式(即 Tuple[Placement] 值)。对于非 Tensor 输出,PlacementType 应为 None。请注意,唯一的例外是没有传入 DTensor 参数的情况。在这种情况下,即使 out_placements 不为 None,结果函数也应忽略所需的放置方式,因为该函数没有使用 DTensor 运行。
2410
+
2411
+ in_placements (Tuple[PlacementType, …], 可选) – func 的展平输入中 DTensor 所需的放置方式。如果指定了 in_placements,local_map() 将检查每个 DTensor 参数的放置方式是否与所需的放置方式相同。如果放置方式不同并且 redistribute_inputs 为 False,则会引发异常。如果 redistribute_inputs 为 True,则在将其本地张量传递给 func 之前,将首先把参数重新分配为所需的分片放置方式。唯一的例外是当所需的放置方式不为 None 且参数是 torch.Tensor 时。在这种情况下,将跳过放置方式检查,并直接将参数传递给 func。如果 in_placements 为 None,则不会执行任何放置方式检查。默认: None
2412
+
2413
+ in_grad_placements (Tuple[PlacementType, …], 可选) – 与展平输入 DTensor 对应的 DTensor 梯度的放置提示。此参数是用户可以提供给 to_local() 的提示,以防本地张量输入的梯度布局与其 DTensor 输入布局不匹配。如果未指定,我们将假设本地张量输入的梯度布局保持与原始 DTensor 输入相同,并将其用于梯度计算。默认: None。
2414
+
2415
+ device_mesh (DeviceMesh, 可选) – 输出 DTensor 放置的设备网格。如果未指定,这将从第一个输入 DTensor 的设备网格中推断出来。默认: None。
2416
+
2417
+ redistribute_inputs (bool, 可选) –指示当输入 DTensor 的放置方式与所需的输入放置方式不同时是否对其进行重新分片的布尔值。如果此值为 False 并且某个 DTensor 输入具有不同的放置方式,则会引发异常。默认: False。
2418
+
2419
+ 一个将 func 应用于输入 DTensor 的每个本地分片并返回由 func 的返回值构建的 DTensor 的可调用对象。
2420
+
2421
+ AssertionError – 对于任何非 DTensor 的输出,我们要求其在 out_placements 中对应的输出放置方式为 None。如果不是这种情况,将引发 AssertionError。
2422
+
2423
+ ValueError – 如果 redistribute_inputs=False,但输入 DTensor 根据 in_placements 需要进行重新分配。
2424
+
2425
+ 此 API 目前是实验性的,可能会发生变化。
2426
+
2427
+ register_sharding() 是一个实验性 API,允许用户在张量输入和输出为 DTensor 时为算子注册分片策略。它可能在以下情况非常有用:(1)当算子不存在默认的分片策略时,例如当算子是 DTensor 尚不支持的自定义算子时;(2)当用户希望覆盖现有算子的默认分片策略时。
2428
+
2429
+ op (Union[OpOverload, List[OpOverload]]) – 用于注册自定义分片函数的一个算子或算子列表。
2430
+
2431
+ 一个函数装饰器,可用于包装一个定义了 op 中指定的算子分片策略的函数。定义的分片策略将被注册到 DTensor,并且如果 DTensor 已经实现了该算子,它将覆盖默认的分片策略。自定义分片函数接收与原始算子相同的输入(除非某个参数是 torch.Tensor,它将被替换为 DTensor 内部使用的类张量对象)。该函数应返回一个由 2 元组组成的序列,每个元组指定可接受的输出放置方式及其对应的输入放置方式。
2432
+
2433
+ 此 API 目前是实验性的,可能会发生变化
2434
+
2435
+ ## FullyShardedDataParallel#
2436
+
2437
+ **URL:** https://pytorch.org/docs/stable/fsdp.html
2438
+
2439
+ **目录:**
2440
+ - FullyShardedDataParallel#
2441
+
2442
+ 创建时间:2022年2月2日 | 最后更新时间:2025年6月11日
2443
+
2444
+ 一个用于在数据并行工作进程间对模块参数进行分片的包装器。
2445
+
2446
+ 这受到了 Xu 等人以及 DeepSpeed 的 ZeRO Stage 3 的启发。FullyShardedDataParallel 通常被简称为 FSDP。
2447
+
2448
+ 使用 FSDP 涉及包装你的模块,然后在之后初始化优化器。这是必需的,因为 FSDP 会更改参数变量。
2449
+
2450
+ 在设置 FSDP 时,你需要考虑目标 CUDA 设备。如果设备具有 ID (dev_id),你有三个选项:
2451
+
2452
+ 将该模块放置在该设备上
2453
+
2454
+ 使用 `torch.cuda.set_device(dev_id)` 设置设备
2455
+
2456
+ 将 dev_id 传递给 `device_id` 构造函数参数。
2457
+
2458
+ 这确保了 FSDP 实例的计算设备是目标设备。对于选项 1 和 3,FSDP 初始化始终在 GPU 上进行。对于选项 2,FSDP 初始化在模块的当前设备上进行,该设备可能是 CPU。
2459
+
2460
+ 如果你正在使用 `sync_module_states=True` 标志,你需要确保模块在 GPU 上,或者使用 `device_id` 参数指定一个 CUDA 设备,以便 FSDP 在 FSDP 构造函数中将模块移动到该设备。这是必要的,因为 `sync_module_states=True` 需要 GPU 通信。
2461
+
2462
+ FSDP 还负责将传递给 forward 方法的输入张量移动到 GPU 计算设备,因此你不需要手动将它们从 CPU 移动。
2463
+
2464
+ 对于 `use_orig_params=True`,与 `ShardingStrategy.FULL_SHARD` 不同,`ShardingStrategy.SHARD_GRAD_OP` 在前向传播后会暴露未分片的参数,而不是分片后的参数。如果你想检查梯度,可以使用 `summon_full_params` 方法并配合 `with_grads=True`。
2465
+
2466
+ 在 `limit_all_gathers=True` 的情况下,你可能会在 FSDP 前向传播之前看到一个间隙,此时 CPU 线程没有发出任何内核。这是有意为之,表明速率限制器正在起作用。以这种方式同步 CPU 线程可以防止为后续的全局聚合过度分配内存,并且它实际上不会延迟 GPU 内核的执行。
2467
+
2468
+ 由于自动微分相关的原因,FSDP 在前向和反向计算期间会用 `torch.Tensor` 视图替换被管理模块的参数。如果你的模块的 forward 依赖于保存的参数引用,而不是在每次迭代时重新获取引用,那么它将看不到 FSDP 新创建的视图,并且自动微分将无法正常工作。
2469
+
2470
+ 最后,当使用 `sharding_strategy=ShardingStrategy.HYBRID_SHARD` 且分片进程组在节点内、复制进程组在节点间时,对于某些集群设置,设置 `NCCL_CROSS_NIC=1` 可以帮助缩短复制进程组的 all-reduce 时间。
2471
+
2472
+ 使用 FSDP 时有几个限制需要注意:
2473
+
2474
+ 当使用 CPU 卸载时,FSDP 目前不支持在 `no_sync()` 之外进行梯度累积。这是因为 FSDP 使用新规约的梯度,而不是与任何现有的梯度进行累积,这可能导致不正确的结果。
2475
+
2476
+ FSDP 不支持运行包含在 FSDP 实例中的子模块的前向传播。这是因为子模块的参数将被分片,但子模块本身不是 FSDP 实例,因此它的前向传播将无法适当地 all-gather 完整的参数。
2477
+
2478
+ 由于其注册反向钩子的方式,FSDP 不支持二阶反向传播(double backwards)。
2479
+
2480
+ FSDP 在冻结参数时有一些限制。对于 `use_orig_params=False`,每个 FSDP 实例必须管理全部被冻结或全部未被冻结的参数。对于 `use_orig_params=True`,FSDP 支持混合冻结和未冻结的参数,但建议不要这样做,以防止梯度内存使用量高于预期。
2481
+
2482
+ 从 PyTorch 1.12 开始,FSDP 提供了对共享参数的有限支持。如果你的用例需要增强的共享参数支持,请在此 issue 中发帖。
2483
+
2484
+ 你应该避免在不使用 `summon_full_params` 上下文的情况下在前向和反向之间修改参数,因为修改可能不会持久化。
2485
+
2486
+ module (nn.Module) – 这是要用 FSDP 包装的模块。
2487
+
2488
+ process_group (可选[Union[ProcessGroup, Tuple[ProcessGroup, ProcessGroup]]]) – 这是用于对模型进行分片的进程组,因此也是用于 FSDP 的 all-gather 和 reduce-scatter 集合通信的进程组。如果为 None,则 FSDP 使用默认进程组。对于混合分片策略(如 `ShardingStrategy.HYBRID_SHARD`),用户可以传入一个进程组元组,分别表示用于分片和复制的组。如果为 None,则 FSDP 为用户构造进程组,以在节点内进行分片并在节点间进行复制。(默认: None)
2489
+
2490
+ sharding_strategy (可选[ShardingStrategy]) – 这用于配置分片策略,该策略可以在节省内存和通信开销之间进行权衡。详情请参见 `ShardingStrategy`。(默认: FULL_SHARD)
2491
+
2492
+ cpu_offload (可选[CPUOffload]) – 这用于配置 CPU 卸载。如果设置为 None,则不会发生 CPU 卸载。详情请参见 `CPUOffload`。(默认: None)
2493
+
2494
+ auto_wrap_policy (可选[Union[Callable[[nn.Module, bool, int], bool], ModuleWrapPolicy, CustomPolicy]]) – 这指定了将 FSDP 应用于 `module` 的子模块的策略,这对于通信和计算的重叠是必需的,因此会影响性能。如果为 None,则 FSDP 仅应用于 `module`,用户应自行手动将 FSDP 应用于父模块(自底向上进行)。为了方便起见,它直接接受 `ModuleWrapPolicy`,允许用户指定要包装的模块类(例如 transformer 块)。否则,它应该是一个接受三个参数 `module: nn.Module`、`recurse: bool` 和 `nonwrapped_numel: int` 的可调用对象,并且如果 `recurse=False`,应返回一个布尔值指定传入的模块是否应应用 FSDP,或者如果 `recurse=True`,指定是否应继续遍历到该模块的子树中。用户可以向该可调用对象添加额外的参数。`torch.distributed.fsdp.wrap.py` 中的 `size_based_auto_wrap_policy` 给出了一个示例可调用对象,如果模块子树中的参数超过 1 亿(100M)个元素,则将 FSDP 应用于该模块。我们建议在应用 FSDP 后打印模型并根据需要进行调整。 示例: >>> def custom_auto_wrap_policy( >>> module: nn.Module, >>> recurse: bool, >>> nonwrapped_numel: int, >>> # Additional custom arguments >>> min_num_params: int = int(1e8), >>> ) -> bool: >>> return nonwrapped_numel >= min_num_params >>> # Configure a custom `min_num_params` >>> my_auto_wrap_policy = functools.partial(custom_auto_wrap_policy, min_num_params=int(1e5))
2495
+
2496
+ 这指定了将 FSDP 应用于 `module` 的子模块的策略,这对于通信和计算的重叠是必需的,因此会影响性能。如果为 None,则 FSDP 仅应用于 `module`,用户应自行手动将 FSDP 应用于父模块(自底向上进行)。为了方便起见,它直接接受 `ModuleWrapPolicy`,允许用户指定要包装的模块类(例如 transformer 块)。否则,它应该是一个接受三个参数 `module: nn.Module`、`recurse: bool` 和 `nonwrapped_numel: int` 的可调用对象,并且如果 `recurse=False`,应返回一个布尔值指定传入的模块是否应应用 FSDP,或者如果 `recurse=True`,指定是否应继续遍历到该模块的子树中。用户可以向该可调用对象添加额外的参数。`torch.distributed.fsdp.wrap.py` 中的 `size_based_auto_wrap_policy` 给出了一个示例可调用对象,如果模块子树中的参数超过 1 亿(100M)个元素,则将 FSDP 应用于该模块。我们建议在应用 FSDP 后打印模型并根据需要进行调整。
2497
+
2498
+ backward_prefetch (可选[BackwardPrefetch]) – 这用于配置 all-gather 的显式反向预取。如果为 None,则 FSDP 不进行反向预取,并且在反向传播中没有通信和计算的重叠。详情请参见 `BackwardPrefetch`。(默认: BACKWARD_PRE)
2499
+
2500
+ mixed_precision (可选[MixedPrecision]) – 这用于配置 FSDP 的原生混合精度。如果设置为 None,则不使用混合精度。否则,可以设置参数、缓冲区和梯度规约的数据类型。详情请参见 `MixedPrecision`。(默认: None)
2501
+
2502
+ ignored_modules (可选[Iterable[torch.nn.Module]]) – 其自身参数以及子模块的参数和缓冲区将被此实例忽略的模块。`ignored_modules` 中直接包含的模块都不应该是 `FullyShardedDataParallel` 实例,并且如果已经是构造好的 `FullyShardedDataParallel` 实例的任何子模块嵌套在此实例下,也不会被忽略。当使用 `auto_wrap_policy` 或者参数的分片不受 FSDP 管理时,此参数可用于在模块粒度上避免对特定参数进行分片。(默认: None)
2503
+
2504
+ param_init_fn (可选[Callable[[nn.Module], None]]) – 一个 `Callable[torch.nn.Module] -> None`,指定了当前位于 meta 设备上的模块应如何初始化到实际设备上。从 v1.12 开始,FSDP 通过 `is_meta` 检测在 meta 设备上具有参数或缓冲区的模块,并在指定时应用 `param_init_fn`,否则调用 `nn.Module.reset_parameters()`。对于这两种情况,实现应该只初始化该模块的参数/缓冲区,而不是其子模块的参数/缓冲区。这是为了避免重复初始化。此外,FSDP 还通过 torchdistX 的 (pytorch/torchdistX) `deferred_init()` API 支持延迟初始化,其中延迟的模块通过调用 `param_init_fn`(如果指定)或 torchdistX 的默认 `materialize_module()` 进行初始化。如果指定了 `param_init_fn`,则它将应用于所有 meta 设备模块,这意味着它可能应该根据模块类型进行分支处理。FSDP 在参数展平和分片之前调用初始化函数。 示例: >>> module = MyModule(device="meta") >>> def my_init_fn(module: nn.Module): >>> # E.g. initialize depending on the module type >>> ... >>> fsdp_model = FSDP(module, param_init_fn=my_init_fn, auto_wrap_policy=size_based_auto_wrap_policy) >>> print(next(fsdp_model.parameters()).device) # current CUDA device >>> # With torchdistX >>> module = deferred_init.deferred_init(MyModule, device="cuda") >>> # Will initialize via deferred_init.materialize_module(). >>> fsdp_model = FSDP(module, auto_wrap_policy=size_based_auto_wrap_policy)
2505
+
2506
+ 一个 `Callable[torch.nn.Module] -> None`,指定了当前位于 meta 设备上的模块应如何初始化到实际设备上。从 v1.12 开始,FSDP 通过 `is_meta` 检测在 meta 设备上具有参数或缓冲区的模块,并在指定时应用 `param_init_fn`,否则调用 `nn.Module.reset_parameters()`。对于这两种情况,实现应该只初始化该模块的参数/缓冲区,而不是其子模块的参数/缓冲区。这是为了避免重复初始化。此外,FSDP 还通过 torchdistX 的 (pytorch/torchdistX) `deferred_init()` API 支持延迟初始化,其中延迟的模块通过调用 `param_init_fn`(如果指定)或 torchdistX 的默认 `materialize_module()` 进行初始化。如果指定了 `param_init_fn`,则它将应用于所有 meta 设备模块,这意味着它可能应该根据模块类型进行分支处理。FSDP 在参数展平和分片之前调用初始化函数。
2507
+
2508
+ device_id (可选[Union[int, torch.device]]) – 一个 int 或 `torch.device`,指示进行 FSDP 初始化(包括模块初始化(如果需要)和参数分片)所在的 CUDA 设备。如果模块在 CPU 上,应指定此项以提高初始化速度。如果设置了默认 CUDA 设备(例如通过 `torch.cuda.set_device`),则用户可以将 `torch.cuda.current_device` 传递给此参数。(默认: None)
2509
+
2510
+ sync_module_states (bool) – 如果为 True,则每个 FSDP 模块将从 rank 0 广播模块参数和缓冲区,以确保它们在各 rank 之间复制(这会给此构造函数增加通信开销)。这有助于通过 `load_state_dict` 以节省内存的方式加载 `state_dict` 检查点。有关此示例,请参见 `FullStateDictConfig`。(默认: False)
2511
+
2512
+ forward_prefetch (bool) – 如果为 True,则 FSDP 会在当前前向计算之前显式预取下一次前向传播的 all-gather。这仅对 CPU 密集型工作负载有用,在这种情况下,提前发出下一次 all-gather 可能会改善重叠。这仅适用于静态图模型,因为预取遵循第一次迭代的执行顺序。(默认: False)
2513
+
2514
+ limit_all_gathers (bool) – 如果为 True,则 FSDP 会显式同步 CPU 线程,以确保仅由两个连续的 FSDP 实例(当前正在运行计算的实例和正在预取 all-gather 的下一个实例)占用 GPU 内存。如果为 False,则 FSDP 允许 CPU 线程发出 all-gather,而无需任何额外的同步。(默认: True) 我们通常将此功能称为“速率限制器”。仅对于内存压力低且受 CPU 限制的特定工作负载,才应将此标志设置为 False,在这种情况下,CPU 线程可以激进地发出所有内核,而无需担心 GPU 内存使用量。
2515
+
2516
+ use_orig_params (bool) – 将此设置为 True 会让 FSDP 使用模块的原始参数。FSDP 通过 `nn.Module.named_parameters()` 将这些原始参数暴露给用户,而不是 FSDP 内部的 `FlatParameter`。这意味着优化器步骤在原始参数上运行,从而允许按原始参数设置超参数。FSDP 保留原始参数变量,并在未分片和分片形式之间操作它们的数据,它们分别始终是底层未分片或分片 `FlatParameter` 的视图。按照当前的算法,分片形式始终是一维的,丢失了原始的张量结构。对于给定的 rank,一个原始参数可能包含其全部、部分或不含任何数据。在不包含数据的情况下,它的数据将类似于一个大小为 0 的空张量。用户不应编写依赖于给定原始参数在其分片形式中存在什么数据的程序。使用 `torch.compile()` 需要将其设置为 True。将此设置为 False 会通过 `nn.Module.named_parameters()` 将 FSDP 内部的 `FlatParameter` 暴露给用户。(默认: False)
2517
+
2518
+ ignored_states (可选[Iterable[torch.nn.Parameter]], 可选[Iterable[torch.nn.Module]]) – 不由此 FSDP 实例管理的被忽略参数或模块,这意味着这些参数不会被分片,并且它们的梯度不会在各 rank 之间进行规约。此参数与现有的 `ignored_modules` 参数统一,我们可能很快弃用 `ignored_modules`。为了向后兼容,我们保留了 `ignored_states` 和 `ignored_modules`,但 FSDP 只允许将其中之一指定为非 None。
2519
+
2520
+ device_mesh (可选[DeviceMesh]) – `DeviceMesh` 可用作 `process_group` 的替代方法。当传递 `device_mesh` 时,FSDP 将使用底层进程组进行 all-gather 和 reduce-scatter 集合通信。因此,这两个参数必须是互斥的。对于混合分片策略(如 `ShardingStrategy.HYBRID_SHARD`),用户可以传入一个二维 `DeviceMesh` 而不是进程组元组。对于二维 FSDP + TP,用户需要传入 `device_mesh` 而不是 `process_group`。有关 `DeviceMesh` 的更多信息,请访问:https://pytorch.org/tutorials/recipes/distributed_device_mesh.html
2521
+
2522
+ 将 `fn` 递归地应用于每个子模块(由 `.children()` 返回)以及自身。
2523
+
2524
+ 典型用途包括初始化模型的参数(另见 `torch.nn.init`)。
2525
+
2526
+ 与 `torch.nn.Module.apply` 相比,此版本在应用 `fn` 之前会额外收集完整的参数。不应在另一个 `summon_full_params` 上下文中调用此方法。
2527
+
2528
+ fn (Module -> None) – 应用于每个子模块的函数
2529
+
2530
+ 检查此实例是否为根 FSDP 模块。
2531
+
2532
+ 裁剪所有参数的梯度范数。
2533
+
2534
+ 该范数是将在所有参数的梯度(视为单个向量)上计算的,并且梯度会被原地修改。
2535
+
2536
+ max_norm (float or int) – 梯度的最大范数
2537
+
2538
+ norm_type (float or int) – 使用的 p 范数的类型。可以为 'inf' 表示无穷范数。
2539
+
2540
+ 参数的总范数(视为单个向量)。
2541
+
2542
+ 如果每个 FSDP 实例都使用 `NO_SHARD`,这意味着没有梯度在各 rank 之间进行分片,那么你可以直接使用 `torch.nn.utils.clip_grad_norm_()`。
2543
+
2544
+ 如果至少有一些 FSDP 实例使用了分片策略(即非 `NO_SHARD` 的策略),那么你应该使用此方法而不是 `torch.nn.utils.clip_grad_norm_()`,因为此方法处理了梯度在各 rank 之间分片的事实。
2545
+
2546
+ 返回的总范数将具有根据 PyTorch 的类型提升语义定义的所有参数/梯度中“最大”的数据类型。例如,如果所有参数/梯度都使用低精度数据类型,则返回范数的数据类型将是该低精度数据类型,但如果存在至少一个使用 FP32 的参数/梯度,则返回范数的数据类型将为 FP32。
2547
+
2548
+ 由于使用集合通信,因此需要在所有 rank 上调用此方法。
2549
+
2550
+ 展平分片的优化器状态字典(state-dict)。
2551
+
2552
+ 该 API 类似于 `shard_full_optim_state_dict()`。唯一的区别在于,输入的 `sharded_optim_state_dict` 应从 `sharded_optim_state_dict()` 返回。因此,每个 rank 上将会有 all-gather 调用来收集 `ShardedTensor`。
2553
+
2554
+ sharded_optim_state_dict (Dict[str, Any]) – 对应于未展平参数并保存分片优化器状态的优化器状态字典。
2555
+
2556
+ model (torch.nn.Module) – 参考 `shard_full_optim_state_dict()`。
2557
+
2558
+ optim (torch.optim.Optimizer) – 针对 `model` 参数的优化器。
2559
+
2560
+ 参考 `shard_full_optim_state_dict()`。
2561
+
2562
+ 运行被包装模块的前向传播,插入 FSDP 特有的前向前和前向后分片逻辑。
2563
+
2564
+ 返回所有嵌套的 FSDP 实例。
2565
+
2566
+ 这可能包含模块本身,并且如果 `root_only=True`,则仅包含 FSDP 根模块。
2567
+
2568
+ module (torch.nn.Module) – 根模块,可能是也可能不是 FSDP 模块。
2569
+
2570
+ root_only (bool) – 是否仅返回 FSDP 根模块。(默认: False)
2571
+
2572
+ 嵌套在输入模块中的 FSDP 模块。
2573
+
2574
+ List[FullyShardedDataParallel]
2575
+
2576
+ 返回完整的优化器状态字典。
2577
+
2578
+ 在 rank 0 上合并完整的优化器状态,并按照 `torch.optim.Optimizer.state_dict()` 的约定将其作为字典返回(即带有键 "state" 和 "param_groups")。包含在 `model` 中的 FSDP 模块中已展平的参数将被映射回其未展平的参数。
2579
+
2580
+ 由于使用集合通信,因此需要在所有 rank 上调用此方法。但是,如果 `rank0_only=True`,则仅在 rank 0 上填充状态字典,所有其他 rank 返回一个空字典。
2581
+
2582
+ 与 `torch.optim.Optimizer.state_dict()` 不同,此方法使用完整的参数名称作为键,而不是参数 ID。
2583
+
2584
+ 与 `torch.optim.Optimizer.state_dict()` 类似,优化器状态字典中包含的张量没有被克隆,因此可能会出现别名问题。作为最佳实践,考虑立即保存返回的优化器状态字典,例如使用 `torch.save()`。
2585
+
2586
+ model (torch.nn.Module) – 根模块(可能是也可能不是 `FullyShardedDataParallel` 实例),其参数被传递给优化器 `optim`。
2587
+
2588
+ optim (torch.optim.Optimizer) – 针对 `model` 参数的优化器。
2589
+
2590
+ optim_input (可选[Union[List[Dict[str, Any]], Iterable[torch.nn.Parameter]]]) – 传递给优化器 `optim` 的输入,表示参数组列表或可迭代的参数;如果为 None,则此方法假定输入为 `model.parameters()`。此参数已弃用,无需再传入。(默认: None)
2591
+
2592
+ rank0_only (bool) – 如果为 True,则仅在 rank 0 上保存填充的字典;如果为 False,则在所有 rank 上保存。(默认: True)
2593
+
2594
+ group (dist.ProcessGroup) – 模型的进程组,如果使用默认进程组则为 None。(默认: None)
2595
+
2596
+ 一个包含 `model` 原始未展平参数的优化器状态的字典,并且遵循 `torch.optim.Optimizer.state_dict()` 的约定包含键 “state” 和 “param_groups”。如果 `rank0_only=True`,则非零 rank 返回一个空字典。
2597
+
2598
+ 获取根植于 `module` 的 FSDP 模块的 `state_dict_type` 及相应的配置。
2599
+
2600
+ 目标模块不必是 FSDP 模块。
2601
+
2602
+ 一个 `StateDictSettings`,包含当前设置的 `state_dict_type` 和 `state_dict` / `optim_state_dict` 配置。
2603
+
2604
+ 如果不同 FSDP 子模块的 `StateDictSettings` 不同,则抛出 `AssertionError`。
2605
+
2606
+ 返回被包装的模块。
2607
+
2608
+ 返回模块缓冲区的迭代器,生成缓冲区的名称和缓冲区本身。
2609
+
2610
+ 在 `summon_full_params()` 上下文管理器内,拦截缓冲区名称并移除所有出现的 FSDP 特定展平缓冲区前缀。
2611
+
2612
+ Iterator[tuple[str, torch.Tensor]]
2613
+
2614
+ 返回模块参数的迭代器,生成参数的名称和参数本身。
2615
+
2616
+ 在 `summon_full_params()` 上下文管理器内,拦截参数名称并移除所有出现的 FSDP 特定展平参数前缀。
2617
+
2618
+ Iterator[tuple[str, torch.nn.parameter.Parameter]]
2619
+
2620
+ 禁用 FSDP 实例之间的梯度同步。
2621
+
2622
+ 在此上下文中,梯度将在模块变量中累积,稍后在退出上下文后的第一次前向-反向传播中进行同步。这仅应在根 FSDP 实例上使用,并将递归应用于所有子 FSDP 实例。
2623
+
2624
+ 这可能会导致更高的内存使用量,因为 FSDP 会累积完整的模型梯度(而不是梯度分片),直到最终同步。
2625
+
2626
+ 当与 CPU 卸载一起使用时,在上下文管理器内梯度不会被卸载到 CPU。相反,它们只会在最终同步之后被卸载。
2627
+
2628
+ 转换对应于分片模型的优化器的 state-dict。
2629
+
2630
+ 给定的 state-dict 可以转换为三种类型之一:1) 完整的优化器 state_dict,2) 分片的优化器 state_dict,3) 本地优化器 state_dict。
2631
+
2632
+ 对于完整的优化器 state_dict,所有状态都是未展平且未分片的。可以通过 `state_dict_type()` 指定仅 rank0 和仅 CPU,以避免 OOM。
2633
+
2634
+ 对于分片的优化器 state_dict,所有状态都是未展平但已分片的。可以通过 `state_dict_type()` 指定仅 CPU 以进一步节省内存。
2635
+
2636
+ 对于本地 state_dict,不会执行任何转换。但状态将从 `nn.Tensor` 转换为 `ShardedTensor` 以表示其分片性质(目前尚不支持)。
2637
+
2638
+ model (torch.nn.Module) – 根模块(可能是也可能不是 `FullyShardedDataParallel` 实例),其参数被传递给优化器 `optim`。
2639
+
2640
+ optim (torch.optim.Optimizer) – 针对 `model` 参数的优化器。
2641
+
2642
+ optim_state_dict (Dict[str, Any]) – 要转换的目标优化器 state_dict。如果值为 None,将使用 `optim.state_dict()`。(默认: None)
2643
+
2644
+ group (dist.ProcessGroup) – 模型的进程组,参数在该组上进行分片;如果使用默认进程组,则为 None。(默认: None)
2645
+
2646
+ 一个包含 `model` 的优化器状态的字典。优化器状态的分片基于 `state_dict_type`。
2647
+
2648
+ 转换优化器状态字典,以便将其加载到与 FSDP 模型关联的优化器中。
2649
+
2650
+ 给定一个通过 `optim_state_dict()` 转换的 `optim_state_dict`,它将被转换为可以加载到 `optim`(即 `model` 的优化器)中的展平优化器状态字典。`model` 必须通过 `FullyShardedDataParallel` 进行分片。
2651
+
2652
+ model (torch.nn.Module) – 根模块(可能是也可能不是 `FullyShardedDataParallel` 实例),其参数被传递给优化器 `optim`。
2653
+
2654
+ optim (torch.optim.Optimizer) – 针对 `model` 参数的优化器。
2655
+
2656
+ optim_state_dict (Dict[str, Any]) – 要加载的优化器状态。
2657
+
2658
+ is_named_optimizer (bool) – 该优化器是否为 `NamedOptimizer` 或 `KeyedOptimizer`。仅当 `optim` 是 TorchRec 的 `KeyedOptimizer` 或 `torch.distributed` 的 `NamedOptimizer` 时才设置为 True。
2659
+
2660
+ load_directly (bool) – 如果设置为 True,此 API 将在返回结果之前调用 `optim.load_state_dict(result)`。否则,用户需自行负责调用 `optim.load_state_dict()` (默认: False)
2661
+
2662
+ group (dist.ProcessGroup) – 模型的进程组,参数在该组上进行分片;如果使用默认进程组,则为 None。(默认: None)
2663
+
2664
+ 注册通信钩子。
2665
+
2666
+ 这是一个增强功能,为用户提供了灵活的钩子,使他们可以指定 FSDP 如何在多个工作进程之间聚合梯度。此钩子可用于实现多种算法,如 GossipGrad 和梯度压缩,这些算法涉及在使用 `FullyShardedDataParallel` 进行训练时用于参数同步的不同通信策略。
2667
+
2668
+ FSDP 通信钩子应在运行初始前向传播之前注册,并且仅能注册一次。
2669
+
2670
+ state (object) – 传递给钩子以在训练过程中维护任何状态信息。示例包括梯度压缩中的误差反馈、GossipGrad 中下一个要与之通信的对等节点等。它由每个工作进程本地存储,并由该工作进程上的所有梯度张量共享。
2671
+
2672
+ 传递给钩子以在训练过程中维护任何状态信息。示例包括梯度压缩中的误差反馈、GossipGrad 中下一个要与之通信的对等节点等。它由每个工作进程本地存储,并由该工作进程上的所有梯度张量共享。
2673
+
2674
+ hook (Callable) – 可调用对象,具有以下签名之一:1) `hook: Callable[torch.Tensor] -> None`:此函数接收一个 Python 张量,它表示相对于此 FSDP 单元所包装的模型对应的所有变量(未被其他 FSDP 子单元包装的)的完整的、展平的、未分片的梯度。然后它执行所有必要的处理并返回 None;2) `hook: Callable[torch.Tensor, torch.Tensor] -> None`:此函数接收两个 Python 张量,第一个表示相对于此 FSDP 单元所包装的模型对应的所有变量(未被其他 FSDP 子单元包装的)的完整的、展平的、未分片的梯度。后者表示一个预先设定大小的张量,用于在规约后存储一块分片的梯度。在这两种情况下,可调用对象执行所有必要的处理并返回 None。签名为 1 的可调用对象预期用于处理 `NO_SHARD` 情况下的梯度通信。签名为 2 的可调用对象预期用于处理分片情况下的梯度通信。
2675
+
2676
+ 重新键控(Re-keys)优化器状态字典 `optim_state_dict`,以使用 `optim_state_key_type` 指定的键类型。
2677
+
2678
+ 这可用于在具有 FSDP 实例的模型和不具有 FSDP 实例的模型之间的优化器状态字典之间实现兼容性。
2679
+
2680
+ 要将 FSDP 完整优化器状态字典(即来自 `full_optim_state_dict()`)重新键控为使用参数 ID,并使其可加载到非包装的模型中:
2681
+
2682
+ 要将来自非包装模型的普通优化器状态字典重新键控为可加载到包装的模型中:
2683
+
2684
+ 使用 `optim_state_key_type` 指定的参数键重新键控的优化器状态字典。
2685
+
2686
+ 将完整的优化器状态字典从 rank 0 分发到所有其他 rank。
2687
+
2688
+ 在每个 rank 上返回分片的优化器状态字典。返回值与 `shard_full_optim_state_dict()` 相同,在 rank 0 上,第一个参数应该是 `full_optim_state_dict()` 的返回值。
2689
+
2690
+ `shard_full_optim_state_dict()` 和 `scatter_full_optim_state_dict()` 都可用于获取要加载的分片优化器状态字典。假设完整的优化器状态字典位于 CPU 内存中,前者要求每个 rank 在 CPU 内存中都有完整的字典,每个 rank 独立地对字典进行分片而无需任何通信;而后者仅要求 rank 0 在 CPU 内存中拥有完整字典,rank 0 将每个分片移动到 GPU 内存(用于 NCCL)并将其适当地通信给各个 rank。因此,前者的总体 CPU 内存成本较高,而后者的通信成本较高。
2691
+
2692
+ full_optim_state_dict (可选[Dict[str, Any]]) – 对应于未展平参数并保存完整非分片优化器状态的优化器状态字典(如果在 rank 0 上);在非零 rank 上忽略此参数。
2693
+
2694
+ model (torch.nn.Module) – 根模块(可能是也可能不是 `FullyShardedDataParallel` 实例),其参数对应于 `full_optim_state_dict` 中的优化器状态。
2695
+
2696
+ optim_input (可选[Union[List[Dict[str, Any]], Iterable[torch.nn.Parameter]]]) – 传递给优化器的输入,表示参数组列表或可迭代的参数;如果为 None,则此方法假定输入为 `model.parameters()`。此参数已弃用,无需再传入。(默认: None)
2697
+
2698
+ optim (可选[torch.optim.Optimizer]) – 将加载由此方法返回的状态字典的优化器。这是比 `optim_input` 更推荐使用的参数。(默认: None)
2699
+
2700
+ group (dist.ProcessGroup) – 模型的进程组,如果使用默认进程组则为 None。(默认: None)
2701
+
2702
+ 现在重新映射为展平参数(而不是未展平参数)并且仅包含此 rank 部分优化器状态的完整优化器状态字典。
2703
+
2704
+ 设置目标模块的所有后代 FSDP 模块的 `state_dict_type`。
2705
+
2706
+ 还可以为模型和优化器的状态字典进行(可选的)配置。目标模块不必是 FSDP 模块。如果目标模块是 FSDP 模块,其 `state_dict_type` 也将被更改。
2707
+
2708
+ 此 API 应仅用于顶层(根)模块。
2709
+
2710
+ 此 API 使用户能够在根 FSDP 模块被另一个 `nn.Module` 包装的情况下,透明地使用传统的 `state_dict` API 来获取模型检查点。例如,以下操作将确保在所有非 FSDP 实例上调用 `state_dict`,同时为 FSDP 调度到 `sharded_state_dict` 实现:
2711
+
2712
+ module (torch.nn.Module) – 根模块。
2713
+
2714
+ state_dict_type (StateDictType) – 要设置的期望的 `state_dict_type`。
2715
+
2716
+ state_dict_config (可选[StateDictConfig]) – 目标 `state_dict_type` 的配置。
2717
+
2718
+ optim_state_dict_config (可选[OptimStateDictConfig]) – 优化器状态字典的配置。
2719
+
2720
+ 一个 `StateDictSettings`,包含模块先前的 `state_dict` 类型和配置。
2721
+
2722
+ 分片完整的优化器状态字典。
2723
+
2724
+ 将 `full_optim_state_dict` 中的状态重新映射为展平参数(而不是未展平参数),并仅限制为此 rank 部分的优化器状态。第一个参数应该是 `full_optim_state_dict()` 的返回值。
2725
+
2726
+ `shard_full_optim_state_dict()` 和 `scatter_full_optim_state_dict()` 都可用于获取要加载的分片优化器状态字典。假设完整的优化器状态字典位于 CPU 内存中,前者要求每个 rank 在 CPU 内存中都有完整的字典,每个 rank 独立地对字典进行分片而无需任何通信;而后者仅要求 rank 0 在 CPU 内存中拥有完整字典,rank 0 将每个分片移动到 GPU 内存(用于 NCCL)并将其适当地通信给各个 rank。因此,前者的总体 CPU 内存成本较高,而后者的通信成本较高。
2727
+
2728
+ full_optim_state_dict (Dict[str, Any]) – 对应于未展平参数并保存完整非分片优化器状态的优化器状态字典。
2729
+
2730
+ model (torch.nn.Module) – 根模块(可能是也可能不是 `FullyShardedDataParallel` 实例),其参数对应于 `full_optim_state_dict` 中的优化器状态。
2731
+
2732
+ optim_input (可选[Union[List[Dict[str, Any]], Iterable[torch.nn.Parameter]]]) – 传递给优化器的输入,表示参数组列表或可迭代的参数;如果为 None,则此方法假定输入为 `model.parameters()`。此参数已弃用,无需再传入。(默认: None)
2733
+
2734
+ optim (可选[torch.optim.Optimizer]) – 将加载由此方法返回的状态字典的优化器。这是比 `optim_input` 更推荐使用的参数。(默认: None)
2735
+
2736
+ 现在重新映射为展平参数(而不是未展平参数)并且仅包含此 rank 部分优化器状态的完整优化器状态字典。
2737
+
2738
+ 返回分片形式的优化器状态字典。
2739
+
2740
+ 该 API 类似于 `full_optim_state_dict()`,但此 API 将所有非零维度的状态分块为 `ShardedTensor` 以节省内存。仅当模型状态字典是通过 `with state_dict_type(SHARDED_STATE_DICT):` 上下文管理器派生时,才应使用此 API。
2741
+
2742
+ 有关详细用法,请参考 `full_optim_state_dict()`。
2743
+
2744
+ 返回的状态字典包含 `ShardedTensor`,不能直接被常规的 `optim.load_state_dict` 使用。
2745
+
2746
+ 设置目标模块的所有后代 FSDP 模块的 `state_dict_type`。
2747
+
2748
+ 此上下文管理器具有与 `set_state_dict_type()` 相同的功能。有关详细信息,请阅读 `set_state_dict_type()` 的文档。
2749
+
2750
+ module (torch.nn.Module) – 根模块。
2751
+
2752
+ state_dict_type (StateDictType) – 要设置的期望的 `state_dict_type`。
2753
+
2754
+ state_dict_config (可选[StateDictConfig]) – 针对 `state_dict_type` 的模型状态字典配置。
2755
+
2756
+ optim_state_dict_config (可选[OptimStateDictConfig]) – 针对 `state_dict_type` 的优化器状态字典配置。
2757
+
2758
+ 使用此上下文管理器为 FSDP 实例暴露完整参数。
2759
+
2760
+ 在前向/反向传播之后,这对模型获取参数以进行额外处理或检查可能很有用。它可以接受非 FSDP 模块,并将根据 `recurse` 参数为所有包含的 FSDP 模块及其子模块调用全部参数。
2761
+
2762
+ 这可以在内部 FSDP 上使用。
2763
+
2764
+ 这不能在前向或反向传播中使用。前向和反向也不能从此上下文中开始。
2765
+
2766
+ 上下文管理器退出后,参数将恢复为其本地分片,存储行为与前向传播相同。
2767
+
2768
+ 可以修改完整参数,但只有与本地参数分片相对应的部分会在上下文管理器退出后保留(除非 `writeback=False`,在这种情况下更改将被丢弃)。在 FSDP 不对参数进行分片的情况下(目前仅在 `world_size == 1` 或 `NO_SHARD` 配置时),无论 `writeback` 如何,修改都会保留。
2769
+
2770
+ 此方法适用于本身不是 FSDP 但可能包含多个独立 FSDP 单元的模块。在这种情况下,给定的参数将应用于所有包含的 FSDP 单元。
2771
+
2772
+ 请注意,目前不支持 `rank0_only=True` 结合 `writeback=True`,并且会引发错误。这是因为在此上下文中,模型参数的形状在各 rank 之间会不同,向它们写入数据可能会导致退出上下文时各 rank 之间的不一致。
2773
+
2774
+ 请注意,`offload_to_cpu` 和 `rank0_only=False` 将导致位于同一台机器上的 GPU 将完整参数冗余复制到 CPU 内存中,这可能会导致 CPU OOM 的风险。建议使用 `offload_to_cpu` 并结合 `rank0_only=True`。
2775
+
2776
+ recurse (bool, 可选) – 递归调用嵌套 FSDP 实例的所有参数 (默认: True)。
2777
+
2778
+ writeback (bool, 可选) – 如果为 False,则在上下文管理器退出后放弃对参数的修改;禁用此选项可能会稍微提高效率 (默认: True)
2779
+
2780
+ rank0_only (bool, 可选) – 如果为 True,则仅在全局 rank 0 上实例化完整参数。这意味着在此上下文中,只有 rank 0 将拥有完整参数,而其他 rank 将拥有分片参数。请注意,不支持将 `rank0_only=True` 与 `writeback=True` 一起使用,因为在此上下文中,模型参数的形状在各 rank 之间会不同,向它们写入数据可能会导致退出上下文时各 rank 之间的不一致。
2781
+
2782
+ offload_to_cpu (bool, 可选) – 如果为 True,则将完整参数卸载到 CPU。请注意,目前仅当参数被分片时(仅在 `world_size = 1` 或 `NO_SHARD` 配置下不是这种情况)才会发生此卸载。建议使用 `offload_to_cpu` 并结合 `rank0_only=True`,以避免将模型参数的冗余副本卸载到相同的 CPU 内存中。
2783
+
2784
+ with_grads (bool, 可选) – 如果为 True,梯度也会与参数一起取消分片。目前,仅当向 FSDP 构造函数传递 `use_orig_params=True` 且向此方法传递 `offload_to_cpu=False` 时,才支持此选项。(默认: False)
2785
+
2786
+ 这配置了显式的反向预取,它通过在反向传播中实现通信和计算重叠来提高吞吐量,代价是内存使用量略微增加。
2787
+
2788
+ BACKWARD_PRE: 这启用了最大程度的重叠,但内存使用量增加最多。这在当前一组参数的梯度计算之前预取下一组参数。这使下一次 all-gather 和当前的梯度计算重叠,并且在峰值时,它会在内存中保存当前这组参数、下一组参数
2789
+
2790
+ ## 分布式优化器#
2791
+
2792
+ **URL:** https://pytorch.org/docs/stable/distributed.optim.html
2793
+
2794
+ **目录:**
2795
+ - 分布式优化器#
2796
+
2797
+ 创建于:2021年3月1日 | 最后更新于:2025年6月16日
2798
+
2799
+ 当使用 CUDA 张量时,目前不支持分布式优化器
2800
+
2801
+ torch.distributed.optim 公开了 `DistributedOptimizer`,它接收一个远程参数列表(`RRef`),并在参数所在的工作节点上本地运行优化器。分布式优化器可以使用任何本地优化器基类在每个工作节点上应用梯度。
2802
+
2803
+ `DistributedOptimizer` 接收分布在各个工作节点上的参数的远程引用,并在本地为每个参数应用给定的优化器。
2804
+
2805
+ 该类使用 `get_gradients()` 来检索特定参数的梯度。
2806
+
2807
+ 对 `step()` 的并发调用(无论是来自同一客户端还是不同客户端)都将在每个工作节点上被串行化——因为每个工作节点的优化器一次只能处理一组梯度。但是,不能保证完整的“前向-反向-优化器”序列会一次只针对一个客户端执行。这意味着,所应用的梯度可能不对应于在给定工作节点上执行的最新前向传播。此外,跨工作节点也没有保证的顺序。
2808
+
2809
+ `DistributedOptimizer` 默认启用 TorchScript 创建本地优化器,以便在多线程训练(例如分布式模型并行)的情况下,优化器的更新不会被 Python 全局解释器锁(GIL)阻塞。此功能目前已为大多数优化器启用。你也可以按照 PyTorch 教程中的指南,为你自己的自定义优化器启用 TorchScript 支持。
2810
+
2811
+ optimizer_class (optim.Optimizer) – 在每个工作节点上实例化的优化器类。
2812
+
2813
+ params_rref (list[RRef]) – 要优化的本地或远程参数的 RRef 列表。
2814
+
2815
+ args – 传递给每个工作节点上优化器构造函数的参数。
2816
+
2817
+ kwargs – 传递给每个工作节点上优化器构造函数的参数。
2818
+
2819
+ 执行单次优化步骤。
2820
+
2821
+ 这将在每个包含待优化参数的工作节点上调用 `torch.optim.Optimizer.step()`,并阻塞直到所有工作节点返回。提供的 `context_id` 将用于检索包含应应用于参数的梯度的相应上下文。
2822
+
2823
+ context_id – 我们应为其运行优化器步骤的 autograd 上下文 ID。
2824
+
2825
+ 包装任意的 `torch.optim.Optimizer` 并运行 post-local SGD,此优化器在每一步运行本地优化器。在预热阶段之后,它会在应用本地优化器后定期对参数进行平均。
2826
+
2827
+ optim (Optimizer) – 本地优化器。
2828
+
2829
+ averager (ModelAverager) – 用于运行 post-localSGD 算法的模型平均器实例。
2830
+
2831
+ 这与 `torch.optim.Optimizer` 的 `load_state_dict()` 相同,但还会将模型平均器的步数值恢复为提供的 `state_dict` 中保存的值。
2832
+
2833
+ 如果 `state_dict` 中没有 "step" 条目,它将发出警告并将模型平均器的步数初始化为 0。
2834
+
2835
+ 这与 `torch.optim.Optimizer` 的 `state_dict()` 相同,但添加了一个额外的条目,用于将模型平均器的步数记录到检查点中,以确保重新加载时不会再次导致不必要的预热。
2836
+
2837
+ 执行单次优化步骤(参数更新)。
2838
+
2839
+ 包装任意的 `optim.Optimizer` 并将其状态分片到组中的各个 rank 上。
2840
+
2841
+ 共享机制按照 ZeRO 的描述进行。
2842
+
2843
+ 每个 rank 中的本地优化器实例仅负责更新大约 1 / world_size 的参数,因此只需保留 1 / world_size 的优化器状态。在本地更新参数之后,每个 rank 会将其参数广播给所有其他节点,以保持所有模型副本处于相同状态。`ZeroRedundancyOptimizer` 可以与 `torch.nn.parallel.DistributedDataParallel` 结合使用,以降低每个 rank 的峰值内存消耗。
2844
+
2845
+ `ZeroRedundancyOptimizer` 使用排序贪心算法在每个 rank 上打包多个参数。每个参数仅属于单个 rank,并且不会在 rank 之间进行切分。这种划分是任意的,可能不匹配参数的注册或使用顺序。
2846
+
2847
+ params (Iterable) – 包含所有参数的 `torch.Tensor` 或 `dict` 的可迭代对象,这些参数将跨 rank 进行分片。
2848
+
2849
+ optimizer_class (torch.nn.Optimizer) – 本地优化器的类。
2850
+
2851
+ process_group (ProcessGroup, 可选) – `torch.distributed` ProcessGroup(默认值:由 `torch.distributed.init_process_group()` 初始化的 `dist.group.WORLD`)。
2852
+
2853
+ parameters_as_bucket_view (bool, 可选) – 如果为 True,参数被打包到桶中以加速通信,并且 `param.data` 字段指向不同偏移量的桶视图;如果为 False,每个单独的参数将单独通信,并且每个 `params.data` 保持不变(默认值:False)。
2854
+
2855
+ overlap_with_ddp (bool, 可选) – 如果为 True,`step()` 将与 `DistributedDataParallel` 的梯度同步重叠;这要求(1)`optimizer_class` 参数要么是函数式优化器,要么是其具有函数式等价物的优化器,并且(2)注册一个由 `ddp_zero_hook.py` 中的函数之一构建的 DDP 通信钩子;参数被打包到与 `DistributedDataParallel` 中的桶相匹配的桶中,这意味着 `parameters_as_bucket_view` 参数将被忽略。如果为 False,`step()` 将在反向传播之后独立运行(按正常情况)。(默认值:False)
2856
+
2857
+ **defaults – 任何尾随参数,将直接转发给本地优化器。
2858
+
2859
+ 目前,`ZeroRedundancyOptimizer` 要求传入的所有参数都是相同的稠密类型。
2860
+
2861
+ 如果传入 `overlap_with_ddp=True` 请注意以下事项:鉴于目前将 `DistributedDataParallel` 与 `ZeroRedundancyOptimizer` 重叠的实现方式,前两到三次训练迭代不会在优化器步骤中执行参数更新,具体分别取决于 `static_graph=False` 还是 `static_graph=True`。这是因为它需要获取有关 `DistributedDataParallel` 使用的梯度分桶策略的信息,而这些信息在 `static_graph=False` 时直到第二次前向传播才会最终确定,或者在 `static_graph=True` 时直到第三次前向传播才会最终确定。为了对此进行调整,一种选择是添加虚拟输入。
2862
+
2863
+ `ZeroRedundancyOptimizer` 处于实验阶段,可能会发生更改。
2864
+
2865
+ 向 `Optimizer` 的 `param_groups` 添加参数组。
2866
+
2867
+ 这在微调预训练网络时非常有用,因为随着训练的进行,被冻结的层可以变为可训练状态,并添加到 `Optimizer` 中。
2868
+
2869
+ param_group (dict) – 指定要优化的参数以及特定于组的优化选项。
2870
+
2871
+ 此方法处理更新所有分区上的分片,但必须在所有 rank 上调用。如果在部分 rank 上调用,会导致训练挂起,因为通信原语是根据受管参数调用的,并且期望所有 rank 都在同一组参数上参与。
2872
+
2873
+ 将 `state_dict` 列表(每个 rank 一个)合并到目标 rank 上。
2874
+
2875
+ to (int) – 接收优化器状态的 rank(默认值:0)。
2876
+
2877
+ RuntimeError – 如果 `overlap_with_ddp=True` 并且在此 `ZeroRedundancyOptimizer` 实例完全初始化之前调用此方法(一旦 `DistributedDataParallel` 梯度桶被重建,此实例即被完全初始化)。
2878
+
2879
+ 这必须在所有 rank 上调用。
2880
+
2881
+ 返回默认设备。
2882
+
2883
+ 返回 ZeRO join hook。
2884
+
2885
+ 它通过遮蔽优化器步骤中的集合通信来支持在不均匀输入上进行训练。
2886
+
2887
+ 在调用此钩子之前必须正确设置梯度。
2888
+
2889
+ kwargs (dict) – 一个包含在运行时修改 join hook 行为的任何关键字参数的字典;所有共享同一 join 上下文管理器的可加入实例都会为 `kwargs` 转发相同的值。
2890
+
2891
+ 此 hook 不支持任何关键字参数;即 `kwargs` 未被使用。
2892
+
2893
+ 返回进程组。
2894
+
2895
+ 从输入的 `state_dict` 加载与给定 rank 相关的状态,并根据需要更新本地优化器。
2896
+
2897
+ state_dict (dict) – 优化器状态;应该是从调用 `state_dict()` 返回的对象。
2898
+
2899
+ RuntimeError – 如果 `overlap_with_ddp=True` 并且在此 `ZeroRedundancyOptimizer` 实例完全初始化之前调用此方法(一旦 `DistributedDataParallel` 梯度桶被重建,此实例即被完全初始化)。
2900
+
2901
+ 返回此 rank 已知的最后一次全局优化器状态。
2902
+
2903
+ RuntimeError – 如果 `overlap_with_ddp=True` 并且在此 `ZeroRedundancyOptimizer` 实例完全初始化之前调用此方法(一旦 `DistributedDataParallel` 梯度桶被重建,此实例即被完全初始化);或者在调用 `consolidate_state_dict()` 之前调用了此方法。
2904
+
2905
+ 执行单次优化器步骤并跨所有 rank 同步参数。
2906
+
2907
+ closure (Callable) – 一个重新评估模型并返回 loss 的闭包;对于大多数优化器是可选的。
2908
+
2909
+ 可选的 loss,具体取决于底层的本地优化器。
2910
+
2911
+ 任何额外的参数都将原样传递给基础优化器。
2912
+
2913
+ ---
2914
+
2915
+ ## Torch 分布式弹性#
2916
+
2917
+ **URL:** https://pytorch.org/docs/stable/distributed.elastic.html
2918
+
2919
+ **目录:**
2920
+ - Torch 分布式弹性#
2921
+ - 快速开始#
2922
+ - 文档#
2923
+
2924
+ 创建于:2025年6月16日 | 最后更新于:2025年7月25日
2925
+
2926
+ 使分布式 PyTorch 具备容错性和弹性。
2927
+
2928
+ ---
2929
+
2930
+ ## 流水线并行#
2931
+
2932
+ **URL:** https://pytorch.org/docs/stable/distributed.pipelining.html
2933
+
2934
+ **目录:**
2935
+ - 流水线并行#
2936
+ - 为什么需要流水线并行?#
2937
+ - 什么是 torch.distributed.pipelining?#
2938
+ - 步骤 1:构建 PipelineStage#
2939
+ - 步骤 2:使用 PipelineSchedule 执行#
2940
+ - 拆分模型的选项#
2941
+ - 选项 1:手动拆分模型#
2942
+ - 选项 2:自动拆分模型#
2943
+ - Hugging Face 示例#
2944
+ - 技术深入剖析#
2945
+
2946
+ 创建于:2025年6月16日 | 最后更新于:2025年8月13日
2947
+
2948
+ `torch.distributed.pipelining` 目前处于 alpha 阶段,正在开发中。API 可能会发生变化。它是从 PiPPy 项目迁移而来的。
2949
+
2950
+ 流水线并行是深度学习的基本并行方式之一。它允许将模型的执行进行分区,使得多个微批次可以并发执行模型代码的不同部分。流水线并行对于以下情况可能是一种有效的技术:
2951
+
2952
+ 带宽受限的集群
2953
+
2954
+ 大模型推理
2955
+
2956
+ 上述场景都有一个共同点,即每个设备上的计算无法掩盖传统并行的通信开销,例如 FSDP 的权重全收集。
2957
+
2958
+ 虽然在扩展方面很有前景,但流水线通常难以实现,因为除了模型权重之外,它还需要对模型的执行进行分区。对执行进行分区通常需要对模型进行侵入性的代码修改。另一个复杂性来自于在分布式环境中调度微批次时需要考虑数据流依赖性。
2959
+
2960
+ `pipelining` 包提供了一个能够自动执行上述操作的工具包,从而允许在通用模型上轻松实现流水线并行。
2961
+
2962
+ 它由两部分组成:拆分前端和分布式运行时。拆分前端接收你的原始模型代码,将其拆分为“模型分区”,并捕获数据流关系。分布式运行时在不同设备上并行执行流水线阶段,处理微批次拆分、调度、通信和梯度传播等事宜。
2963
+
2964
+ 总体而言,`pipelining` 包提供以下功能:
2965
+
2966
+ 基于简单的规范拆分模型代码。
2967
+
2968
+ 对流水线调度提供丰富的支持,包括 GPipe、1F1B、Interleaved 1F1B 和 Looped BFS,并提供编写自定义调度的基础设施。
2969
+
2970
+ 对跨主机流水线并行提供一流的支持,因为这通常是使用 PP 的场景(通过较慢的互连)。
2971
+
2972
+ 与其他 PyTorch 并行技术(如数据并行(DDP,FSDP)或张量并行)的可组合性。TorchTitan 项目演示了在 Llama 模型上的“3D 并行”应用。
2973
+
2974
+ 在我们可以使用 `PipelineSchedule` 之前,我们需要创建 `PipelineStage` 对象来包装在该阶段中运行的模型部分。`PipelineStage` 负责分配通信缓冲区并创建发送/接收操作以与其对等端通信。它管理中间缓冲区(例如用于尚未被消耗的前向传播输出),并提供为阶段模型运行反向传播的工具。
2975
+
2976
+ `PipelineStage` 需要知道阶段模型的输入和输出形状,以便它可以正确分配通信缓冲区。这些形状必须是静态的,例如在运行时,形状不能在不同的步骤之间发生变化。如果运行时的形状与预期的形状不匹配,将引发 `PipeliningShapeError` 异常。当与其他并行机制组合或应用混合精度时,必须考虑到这些技术,以便 `PipelineStage` 知道阶段模块在运行时输出的正确形状(和 dtype)。
2977
+
2978
+ 用户可以直接构造 `PipelineStage` 实例,方法是传入一个 `nn.Module`,该模块表示应在该阶段上运行的那部分模型。这可能需要修改原始模型代码。请参阅选项 1:手动拆分模型中的示例。
2979
+
2980
+ 或者,拆分前端可以使用图分区自动将你的模型拆分为一系列 `nn.Module`。此技术要求模型可以通过 `torch.Export` 进行跟踪。生成的 `nn.Module` 与其他并行技术的可组合性是实验性的,可能需要一些变通方法。如果用户无法轻易更改模型代码,使用该前端可能会更具吸引力。有关更多信息,请参阅选项 2:自动拆分模型。
2981
+
2982
+ 现在我们可以将 `PipelineStage` 附加到流水线调度上,并使用输入数据运行调度。以下是一个 GPipe 示例:
2983
+
2984
+ 请注意,上述代码需要为每个工作节点启动,因此我们使用启动器服务来启动多个进程:
2985
+
2986
+ 为了直接构造 `PipelineStage`,用户负责提供一个包含相关 `nn.Parameters` 和 `nn.Buffers` 的单一 `nn.Module` 实例,并定义一个执行与该阶段相关的操作的 `forward()` 方法。例如,Torchtitan 中定义的 Transformer 类的精简版本展示了构建易于分区模型的模式。
2987
+
2988
+ 以这种方式定义的模型可以通过以下方式轻松地为每个阶段进行配置:首先初始化整个模型(使用 meta-device 以避免 OOM 错误),删除该阶段不需要的层,然后创建一个包装该模型的 `PipelineStage`。例如:
2989
+
2990
+ 当与其他数据或模型并行技术组合时,如果模型块的输出形状/dtype 会受到影响,也可能需要 `output_args`。
2991
+
2992
+ 如果你有一个完整的模型,并且不想花时间将其修改为一系列“模型分区”,流水线 API 可以为你提供帮助。这是一个简短的示例:
2993
+
2994
+ 如果我们打印该模型,我们可以看到多个层级,这使得手动拆分变得困难:
2995
+
2996
+ 让我们看看流水线 API 是如何工作的:
2997
+
2998
+ 流水线 API 根据给定的 `split_spec` 拆分你的模型,其中 `SplitPoint.BEGINNING` 表示在 forward 函数中某个子模块执行之前添加一个分割点,类似地,`SplitPoint.END` 表示在其执行之后添加一个分割点。
2999
+
3000
+ 如果我们打印 `pipe`,我们可以看到:
3001
+
3002
+ “模型分区”由子模块(`submod_0`,`submod_1`)表示,每个子模块都是用原始模型操作、权重和层次结构重建的。此外,重建了一个“根级别”的 forward 函数来捕获这些分区之间的数据流。这些数据流稍后将以分布式的方式由流水线运行时重放。
3003
+
3004
+ `Pipe` 对象提供了一个用于检索“模型分区”的方法:
3005
+
3006
+ 返回的 `stage_mod` 是一个 `nn.Module`,你可以使用它来创建优化器、保存或加载检查点,或应用其他并行机制。
3007
+
3008
+ `Pipe` 还允许你在给定 `ProcessGroup` 的情况下在设备上创建分布式阶段运行时:
3009
+
3010
+ 或者,如果你想在修改 `stage_mod` 之后再构建阶段运行时,你可以使用函数版本的 `build_stage` API。例如:
3011
+
3012
+ 流水线前端使用跟踪器(`torch.export`)将你的模型捕获到一个单一的图中。如果你的模型无法进行全图捕获,你可以使用下面的手动前端。
3013
+
3014
+ 在最初创建此包的 PiPPy 存储库中,我们保留了基于未修改的 Hugging Face 模型的示例。请参阅 `examples/huggingface` 目录。
3015
+
3016
+ 首先,流水线 API 通过跟踪模型将其转换为有向无环图(DAG)。它使用 `torch.export`——一种 PyTorch 2 全图捕获工具来跟踪模型。
3017
+
3018
+ 然后,它将一个阶段所需的操作和参数组合在一起,重建为一个子模块:`submod_0`,`submod_1`,...
3019
+
3020
+ 与传统的子模块访问方法(如 `Module.children()`)不同,流水线 API 不仅会切断模型的结构,还会切断模型的 forward 函数。
3021
+
3022
+ 这是必要的,因为像 `Module.children()` 这样的模型结构仅在 `Module.__init__()` 期间捕获信息,而不捕获有关 `Module.forward()` 的任何信息。换句话说,`Module.children()` 缺少有关以下对流水线至关重要的方面的信息:
3023
+
3024
+ 子模块在 forward 中的执行顺序
3025
+
3026
+ 子模块之间的激活流
3027
+
3028
+ 子模块之间是否存在函数式运算符(例如,`relu` 或 `add` 操作不会被 `Module.children()` 捕获)。
3029
+
3030
+ 相反,流水线 API 确保真正保留了 forward 行为。它还捕获分区之间的激活流,帮助分布式运行时在无需人工干预的情况下发出正确的发送/接收调用。
3031
+
3032
+ 流水线 API 的另一个灵活性在于,分割点可以位于模型层次结构中的任意级别。在分割的分区中,与该分区相关的原始模型层次结构将为你重建,不产生额外成本。结果,指向子模块或参数的完全限定名(FQN)仍然有效,依赖于 FQN 的服务(如 FSDP、TP 或检查点)可以在你的分区模块上运行,几乎不需要更改代码。
3033
+
3034
+ 你可以通过扩展以下两个类之一来实现自己的流水线调度:
3035
+
3036
+ `PipelineScheduleSingle`
3037
+
3038
+ `PipelineScheduleMulti`
3039
+
3040
+ `PipelineScheduleSingle` 用于每个 rank 仅分配一个阶段的调度。`PipelineScheduleMulti` 用于每个 rank 分配多个阶段的调度。
3041
+
3042
+ 例如,`ScheduleGPipe` 和 `Schedule1F1B` 是 `PipelineScheduleSingle` 的子类。而 `ScheduleInterleaved1F1B`、`ScheduleLoopedBFS`、`ScheduleInterleavedZeroBubble` 和 `ScheduleZBVZeroBubble` 是 `PipelineScheduleMulti` 的子类。
3043
+
3044
+ 你可以使用 `torch._logging` 中的 `TORCH_LOGS` 环境变量开启附加日志记录:
3045
+
3046
+ `TORCH_LOGS=+pp` 将显示 `logging.DEBUG` 消息及高于它的所有级别的消息。
3047
+
3048
+ `TORCH_LOGS=pp` 将显示 `logging.INFO` 消息及更高级别的消息。
3049
+
3050
+ `TORCH_LOGS=-pp` 将显示 `logging.WARNING` 消息及更高级别的消息。
3051
+
3052
+ 下面的这组 API 将你的模型转换为流水线表示。
3053
+
3054
+ 表示子模块执行过程中可以发生分割的点的枚举。 :ivar BEGINNING: 表示在 forward 函数中某个子模块执行之前添加分割点。 :ivar END: 表示在 forward 函数中某个子模块执行之后添加分割点。
3055
+
3056
+ 基于规范拆分模块。
3057
+
3058
+ 有关更多详细信息,请参见 `Pipe`。
3059
+
3060
+ module (Module) – 要拆分的模块。
3061
+
3062
+ mb_args (tuple[Any, ...]) – 微批次形式的示例位置参数输入。
3063
+
3064
+ mb_kwargs (可选[dict[str, Any]]) – 微批次形式的示例关键字参数输入。(默认值:None)
3065
+
3066
+ split_spec (可选[dict[str, torch.distributed.pipelining._IR.SplitPoint]]) – 使用子模块名称作为分割标记的字典。(默认值:None)
3067
+
3068
+ split_policy (可选[Callable[[GraphModule], GraphModule]]) – 用于拆分模块的策略。(默认值:None)
3069
+
3070
+ `Pipe` 类的流水线表示。
3071
+
3072
+ `pipe_split` 是一个特殊的运算符,用于标记模块中阶段之间的边界。它用于将模块划分为多个阶段。如果你标记的模块在 eager 模式下运行,它是一个空操作。
3073
+
3074
+ 上面的示例将被分割成两个阶段。
3075
+
3076
+ 用于指定输入分块的类
3077
+
3078
+ 根据它们各自的分块规范,将给定的 args 和 kwargs 序列分割成若干块。
3079
+
3080
+ args (tuple[Any, ...]) – args 组成的元组
3081
+
3082
+ kwargs (可选[dict[str, Any]]) – kwargs 组成的字典
3083
+
3084
+ chunks (int) – 将 args 和 kwargs 切分成的块数
3085
+
3086
+ args_chunk_spec (可选[tuple[torch.distributed.pipelining.microbatch.TensorChunkSpec, ...]]) – args 的分块规范,形状与 args 相同
3087
+
3088
+ kwargs_chunk_spec (可选[dict[str, torch.distributed.pipelining.microbatch.TensorChunkSpec]]) – kwargs 的分块规范,形状与 kwargs 相同
3089
+
3090
+ 分片后的 args 列表 kwargs_split: 分片后的 kwargs 列表
3091
+
3092
+ 根据分块规范将给定的块列表合并为单个值。
3093
+
3094
+ chunks (list[Any]) – 块组成的列表
3095
+
3096
+ chunk_spec – 各个块的分布规范
3097
+
3098
+ 表示流水线并行设置中一个流水线阶段的类。
3099
+
3100
+ `PipelineStage` 假设模型进行顺序分区,即模型被分成多个块,其中一个块的输出馈送到下一个块的输入中,没有跳跃连接。
3101
+
3102
+ `PipelineStage` 通过将输出从 stage0 线性顺序地传播到 stage1 等,自动执行运行时形状/dtype 推断。要绕过形状推断,请将 `input_args` 和 `output_args` 传递给每个 `PipelineStage` 实例。
3103
+
3104
+ submodule (nn.Module) – 被此阶段包装的 PyTorch 模块。
3105
+
3106
+ stage_index (int) – 此阶段的 ID。
3107
+
3108
+ num_stages (int) – 阶段总数。
3109
+
3110
+ device (torch.device) – 此阶段所在的设备。
3111
+
3112
+ input_args (Union[torch.Tensor, Tuple[torch.tensor]], 可选) – 子模块的输入参数。
3113
+
3114
+ output_args (Union[torch.Tensor, Tuple[torch.tensor]], 可选) – 子模块的输出参数。
3115
+
3116
+ group (dist.ProcessGroup, 可选) – 用于分布式训练的进程组。如果为 None,则使用默认组。
3117
+
3118
+ dw_builder (可选[Callable[[], Callable[..., None]]) – 如果提供,`dw_builder` 将构建一个新的 `dw_runner` 函数,用于在 F、I、W(前向、输入、权重)零气泡调度中执行 W 操作(输入权重)。
3119
+
3120
+ 在给定要被此阶段包装的 `stage_module` 和流水线信息的情况下创建一个流水线阶段。
3121
+
3122
+ stage_module (torch.nn.Module) – 被此阶段包装的模块
3123
+
3124
+ stage_index (int) – 此阶段在流水线中的索引
3125
+
3126
+ pipe_info (PipeInfo) – 有关流水线的信息,可以通过 `pipe.info()` 检索
3127
+
3128
+ device (torch.device) – 此阶段使用的设备
3129
+
3130
+ group (可选[dist.ProcessGroup]) – 此阶段使用的进程组
3131
+
3132
+ 一个可与 `PipelineSchedules` 一起运行的流水线阶段。
3133
+
3134
+ GPipe 调度。将以填充-排空的方式处理所有微批次。
3135
+
3136
+ 1F1B 调度。将在稳定状态下对微批次执行一次前向传播和一次反向传播。
3137
+
3138
+ Interleaved 1F1B 调度。有关详细信息,请参见 https://arxiv.org/pdf/2104.04473。将在稳定状态下对微批次执行一次前向传播和一次反向传播,并支持每个 rank 对应多个阶段。当多个本地阶段的微批次准备就绪时,Interleaved 1F1B 会优先处理较早的微批次(也称为“深度优先”)。
3139
+
3140
+ 此调度与原论文大致相同。区别在于它放宽了对 `num_microbatch % pp_size == 0` 的要求。使用 `flex_pp` 调度,我们将得到 `num_rounds = max(1, n_microbatches // pp_group_size)`,只要 `n_microbatches % num_rounds` 为 0,它就能正常工作。举几个支持的例子:
3141
+
3142
+ pp_group_size = 4,n_microbatches = 10。我们会有 num_rounds = 2 并且 n_microbatches % 2 为 0。
3143
+
3144
+ pp_group_size = 4,n_microbatches = 3。我们会有 num_rounds = 1 并且 n_microbatches % 1 为 0。
3145
+
3146
+ 广度优先流水线并行。有关详细信息,请参见 https://arxiv.org/abs/2211.05953。与 Interleaved 1F1B 类似,Looped BFS 支持每个 rank 对应多个阶段。不同之处在于,当多个本地阶段的微批次准备就绪时,Loops BFS 将优先处理较早的阶段,立即运行所有可用的微批次。
3147
+
3148
+ Interleaved Zero Bubble 调度。有关详细信息,请参见 https://arxiv.org/pdf/2401.10241。将在稳定状态下对微批次的输入执行一次前向传播和一次反向传播,并支持每个 rank 对应多个阶段。使用权重的反向传播来填补流水线气泡。
3149
+
3150
+ 具体而言,这是在实现论文中的 ZB1P 调度。
3151
+
3152
+ Zero Bubble 调度(ZBV 变体)。有关详细信息,请参见 https://arxiv.org/pdf/2401.10241 第 6 节。
3153
+
3154
+ 此调度要求每个 rank 恰好有两个阶段。
3155
+
3156
+ 此调度将在稳定状态下对微批次的输入执行一次前向传播和一次反向传播,并支持每个 rank 对应多个阶段。使用关于权重的反向传播来填补流水线气泡。
3157
+
3158
+ 只有当前向时间 == 输入反向时间 == 权重反向时间时,此 ZB-V 调度才具有“零气泡”属性。在实践中,这对于实际模型不太可能是真实的,因此或者可以为不相等/不平衡的时间实现一个贪心调度器。
3159
+
3160
+ DualPipeV 调度。基于 DeepSeek 在 https://arxiv.org/pdf/2412.19437 中引入的 DualPipe 调度的更高效的调度变体
3161
+
3162
+ 基于 deepseek-ai/DualPipe 的开源代码
3163
+
3164
+ 单阶段调度的基类。实现了 `step` 方法。派生类应该实现 `_step_microbatches`。
3165
+
3166
+ 根据 `scale_grads` 参数(默认为 True),梯度会按 `num_microbatches` 进行缩放。此设置应与你的 `loss_fn` 配置相匹配,它可以是对损失求平均值(`scale_grads=True`)或对损失求和(`scale_grads=False`)。
3167
+
3168
+ 使用整批输入运行一次流水线调度的迭代。会自动将输入切分成微批次,并根据调度实现来处理这些微批次。
3169
+
3170
+ args:模型的位置参数(与非流水线情况一样)。kwargs:模型的关键字参数(与非流水线情况一样)。target:损失函数的目标值。losses:一个用于存储每个微批次损失的列表。
3171
+
3172
+ 多阶段调度的基类。实现了 `step` 方法。
3173
+
3174
+ 根据 `scale_grads` 参数(默认为 True),梯度会按 `num_microbatches` 进行缩放。此设置应与你的 `loss_fn` 配置相匹配,它可以是对损失求平均值(`scale_grads=True`)或对损失求和(`scale_grads=False`)。
3175
+
3176
+ 使用整批输入运行一次流水线调度的迭代。会自动将输入切分成微批次,并根据调度实现来处理这些微批次。
3177
+
3178
+ args:模型的位置参数(与非流水线情况一样)。kwargs:模型的关键字参数(与非流水线情况一样)。target:损失函数的目标值。losses:一个用于存储每个微批次损失的列表。
3179
+
3180
+ ---
3181
+
3182
+ ## 张量并行 - torch.distributed.tensor.parallel#
3183
+
3184
+ **URL:** https://pytorch.org/docs/stable/distributed.tensor.parallel.html
3185
+
3186
+ **目录:**
3187
+ - 张量并行 - torch.distributed.tensor.parallel#
3188
+
3189
+ 创建于:2025年6月13日 | 最后更新于:2025年6月13日
3190
+
3191
+ 张量并行(TP)构建在 PyTorch DistributedTensor (DTensor)[https://github.com/pytorch/pytorch/blob/main/torch/distributed/tensor/README.md] 之上,提供了不同的并行方式:Colwise、Rowwise 和 Sequence Parallelism。
3192
+
3193
+ 张量并行 API 是实验性的,随时可能更改。
3194
+
3195
+ 使用张量并行来并行化你的 `nn.Module` 的入口点是:
3196
+
3197
+ 通过基于用户指定的计划并行化模块或子模块,在 PyTorch 中应用张量并行。
3198
+
3199
+ 我们基于 `parallelize_plan` 并行化模块或子模块。`parallelize_plan` 包含 `ParallelStyle`,指示用户希望如何对模块或子模块进行并行化。
3200
+
3201
+ 用户还可以通过模块的完全限定名(FQN)为每个模块指定不同的并行方式。
3202
+
3203
+ 请注意,`parallelize_module` 仅接受一维的 `DeviceMesh`。如果你有二维或 N 维的 `DeviceMesh`,请首先将 `DeviceMesh` 切分为一维的子 `DeviceMesh`,然后再传递给此 API(即 `device_mesh["tp"]`)。
3204
+
3205
+ module (nn.Module) – 要并行化的模块。
3206
+
3207
+ device_mesh (DeviceMesh, 可选) – 描述 DTensor 的设备网格拓扑的对象。如果未指定,则该调用必须在 `DeviceMesh` 上下文中。
3208
+
3209
+ parallelize_plan (Union[ParallelStyle, Dict[str, ParallelStyle]], 可选) – 用于并行化模块的计划。它可以是包含我们如何为张量并行准备输入/输出的 `ParallelStyle` 对象,也可以是模块 FQN 及其相应 `ParallelStyle` 对象的字典。如果未指定,该调用此刻将不执行任何操作。
3210
+
3211
+ src_data_rank (int, 可选) – 逻辑/全局张量的源数据的 rank,它被 `distribute_tensor()` 用于向其他 rank 分发/广播分片/副本。默认情况下,我们在每个 `DeviceMesh` 维度上使用 `group_rank=0` 作为源数据,以保留单设备语义。如果显式传入 None,`parallelize_module()` 将仅使用其本地数据,而不是尝试通过分发/广播来保留单设备语义。默认:0
3212
+
3213
+ 一个被并行化的 `nn.Module` 对象。
3214
+
3215
+ 对于复杂的模块架构(如 Attention、MLP 层),我们建议将不同的 `ParallelStyle`(即 `ColwiseParallel` 和 `RowwiseParallel`)组合在一起,并作为 `parallelize_plan` 传入,以实现所需的分片计算。
3216
+
3217
+ 张量并行支持以下并行方式:
3218
+
3219
+ 按列方式对兼容的 `nn.Module` 进行分区。目前支持 `nn.Linear` 和 `nn.Embedding`。用户可以将其与 `RowwiseParallel` 组合在一起,实现更复杂模块(即 MLP、Attention)的分片。
3220
+
3221
+ input_layouts (Placement, 可选) – `nn.Module` 输入张量的 DTensor 布局,用于标注输入张量使其成为 DTensor。如果未指定,我们假定输入张量是复制。
3222
+
3223
+ output_layouts (Placement, 可选) – `nn.Module` 输出的 DTensor 布局,用于确保 `nn.Module` 的输出具有用户期望的布局。如果未指定,输出张量会在最后一个维度上进行分片。
3224
+
3225
+ use_local_output (bool, 可选) – 是否为模块输出使用本地 `torch.Tensor` 而不是 DTensor,默认值:True。
3226
+
3227
+ 一个表示 `nn.Module` 的 Colwise 分片的 `ParallelStyle` 对象。
3228
+
3229
+ 默认情况下,如果未指定 `output_layouts`,`ColwiseParallel` 的输出会在最后一个维度上进行分片,如果有运算符需要特定的张量形状(即在配对的 `RowwiseParallel` 之前),请记住,如果输出被分片了,可能需要对运算符进行调整以适应分片后的大小。
3230
+
3231
+ 按行方式对兼容的 `nn.Module` 进行分区。目前支持 `nn.Linear` 和 `nn.Embedding`。用户可以将其与 `ColwiseParallel` 组合在一起,实现更复杂模块(即 MLP、Attention)的分片。
3232
+
3233
+ input_layouts (Placement, 可选) – `nn.Module` 输入张量的 DTensor 布局,用于标注输入张量使其成为 DTensor。如果未指定,我们假定输入张量在最后一个维度上进行分片。
3234
+
3235
+ output_layouts (Placement, 可选) – `nn.Module` 输出的 DTensor 布局,用于确保 `nn.Module` 的输出具有用户期望的布局。如果未指定,输出张量为复制。
3236
+
3237
+ use_local_output (bool, 可选) – 是否为模块输出使用本地 `torch.Tensor` 而不是 DTensor,默认值:True。
3238
+
3239
+ 一个表示 `nn.Module` 的 Rowwise 分片的 `ParallelStyle` 对象。
3240
+
3241
+ SequenceParallel 会复制兼容的 `nn.Module` 参数,并运行输入在序列维度上进行分片后的计算。目前支持 `nn.LayerNorm`、`nn.Dropout` 和 RMSNorm 的 python 实现。
3242
+
3243
+ 此方式实现了论文《Reducing Activation Recomputation in Large Transformer Models》中描述的操作。
3244
+
3245
+ 如果传入此 `nn.Module` 的输入是 `torch.Tensor`,它会假定输入已经在序列维度上进行了分片,并将输入转换为在序列维度上分片的 DTensor。如果传入此 `nn.Module` 的输入已经是 DTensor 但并未在序列维度上分片,它将重新分布输入,以实现在序列维度上分片。
3246
+
3247
+ `nn.Module` 的输出将在序列维度上分片。
3248
+
3249
+ sequence_dim (int, 可选) – `nn.Module` 的输入张量的序列维度,用于标注输入张量使其成为在序列维度上分片的 DTensor,默认值:1。
3250
+
3251
+ use_local_output (bool, 可选) – 是否为模块输出使用本地 `torch.Tensor` 而不是 DTensor,默认值:False。
3252
+
3253
+ 一个表示 `nn.Module` 的 Sequence Parallel 的 `ParallelStyle` 对象。
3254
+
3255
+ `SequenceParallel` 方式假定 `nn.Module` 中有权重时(如 `nn.LayerNorm` 或 RMSNorm)进行全 1 初始化(且它们默认就是全 1 初始化)。如果这些模块上的权重有自定义初始化,则需要在并行化前后广播权重,以确保它们被复制。
3256
+
3257
+ 为了仅使用 DTensor 布局配置 `nn.Module` 的输入和输出并执行必要的布局重分布,而不将模块参数分配为 DTensor,可以在调用 `parallelize_module` 时在 `parallelize_plan` 中使用以下 `ParallelStyle`:
3258
+
3259
+ 根据 `input_layouts` 将 `nn.Module` 的输入张量在运行时转换为 DTensor,并根据 `desired_input_layouts` 执行布局重分布,从而配置 `nn.Module` 的输入。
3260
+
3261
+ input_layouts (Union[Placement, Tuple[可选[Placement]]]) – `nn.Module` 的输入张量的 DTensor 布局,用于将输入张量转换为 DTensor。如果某些输入不是 `torch.Tensor` 或无需转换为 DTensor,则需要指定 None 作为占位符。默认值:None。
3262
+
3263
+ desired_input_layouts (Union[Placement, Tuple[可选[Placement]]]) – `nn.Module` 输入张量的所需 DTensor 布局,用于确保 `nn.Module` 的输入具有所需的 DTensor 布局。此参数需要与 `input_layouts` 具有相同的长度。默认值:None。
3264
+
3265
+ input_kwarg_layouts (Dict[str, Placement]) – `nn.Module` 的输入 kwargs 的 DTensor 布局,用于将输入 kwarg 张量转换为 DTensor。默认值:None
3266
+
3267
+ desired_input_kwarg_layouts – (Dict[str, Placement]):`nn.Module` 的输入 kwargs 的所需 DTensor 布局,用于确保 `nn.Module` 的输入具有所需的 DTensor 布局。默认值:None。
3268
+
3269
+ use_local_output (bool, 可选) – 是否为模块输入使用本地 `torch.Tensor` 而不是 DTensor,默认值:False。
3270
+
3271
+ 一个用于准备 `nn.Module` 输入分片布局的 `ParallelStyle` 对象。
3272
+
3273
+ 根据 `output_layouts` 将 `nn.Module` 的输出张量在运行时转换为 DTensor,并根据 `desired_output_layouts` 执行布局重分布,从而配置 `nn.Module` 的输出。
3274
+
3275
+ output_layouts (Union[Placement, Tuple[Placement]]) – `nn.Module` 的输出张量的 DTensor 布局,如果它们是 `torch.Tensor`,则用于将其转换为 DTensor。如果某些输出不是 `torch.Tensor` 或无需转换为 DTensor,则需要指定 None 作为占位符。
3276
+
3277
+ desired_output_layouts (Union[Placement, Tuple[Placement]]) – `nn.Module` 的输出张量的所需 DTensor 布局,用于确保 `nn.Module` 的输出具有所需的 DTensor 布局。
3278
+
3279
+ use_local_output (bool, 可选) – 是否为模块输出使用本地 `torch.Tensor` 而不是 DTensor,默认值:True。
3280
+
3281
+ 一个用于准备 `nn.Module` 输出分片布局的 `ParallelStyle` 对象。
3282
+
3283
+ 根据 `input_layouts`(和 `output_layouts`)将 `nn.Module` 的输入张量(以及相应的输出张量)在运行时转换为 DTensor,并根据 `desired_input_layouts`(以及 `desired_output_layouts`)执行布局重分布,从而配置 `nn.Module` 的输入(和输出)。这是 `PrepareModuleInput` 和 `PrepareModuleOutput` 的组合。
3284
+
3285
+ input_layouts (Union[Placement, Tuple[可选[Placement]]]) – `nn.Module` 的输入张量的 DTensor 布局,用于将输入张量转换为 DTensor。如果某些输入不是 `torch.Tensor` 或无需转换为 DTensor,则需要指定 None 作为占位符。默认值:None。
3286
+
3287
+ desired_input_layouts (Union[Placement, Tuple[可选[Placement]]]) – `nn.Module` 输入张量的所需 DTensor 布局,用于确保 `nn.Module` 的输入具有所需的 DTensor 布局。此参数需要与 `input_layouts` 具有相同的长度。默认值:None。
3288
+
3289
+ input_kwarg_layouts (Dict[str, Placement]) – `nn.Module` 的输入 kwargs 的 DTensor 布局,用于将输入 kwarg 张量转换为 DTensor。默认值:None
3290
+
3291
+ desired_input_kwarg_layouts – (Dict[str, Placement]):`nn.Module` 的输入 kwargs 的所需 DTensor 布局,用于确保 `nn.Module` 的输入具有所需的 DTensor 布局。默认值:None。
3292
+
3293
+ use_local_input (bool, 可选) – 是否为模块输入使用本地 `torch.Tensor` 而不是 DTensor,默认值:False。
3294
+
3295
+ output_layouts (Union[Placement, Tuple[Placement]]) – `nn.Module` 的输出张量的 DTensor 布局,如果它们是 `torch.Tensor`,则用于将其转换为 DTensor。如果某些输出不是 `torch.Tensor` 或无需转换为 DTensor,则需要指定 None 作为占位符。
3296
+
3297
+ desired_output_layouts (Union[Placement, Tuple[Placement]]) – `nn.Module` 的输出张量的所需 DTensor 布局,用于确保 `nn.Module` 的输出具有