diffusers-workflow 0.4.0__py3-none-any.whl
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.
- diffusers_workflow-0.4.0.dist-info/METADATA +318 -0
- diffusers_workflow-0.4.0.dist-info/RECORD +260 -0
- diffusers_workflow-0.4.0.dist-info/WHEEL +5 -0
- diffusers_workflow-0.4.0.dist-info/entry_points.txt +7 -0
- diffusers_workflow-0.4.0.dist-info/licenses/LICENSE +201 -0
- diffusers_workflow-0.4.0.dist-info/top_level.txt +2 -0
- dw/__init__.py +440 -0
- dw/adapter_compatibility.py +226 -0
- dw/arguments.py +1231 -0
- dw/assessment_rules.py +159 -0
- dw/assets.py +130 -0
- dw/cache_blocks.json +16 -0
- dw/cache_blocks.py +146 -0
- dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
- dw/content_types.py +150 -0
- dw/dissolve_frame_errors.py +121 -0
- dw/docs/ACCELERATION.md +352 -0
- dw/docs/AGENT_LOOP.md +95 -0
- dw/docs/DEPENDENCIES.md +91 -0
- dw/docs/IP_ADAPTER.md +109 -0
- dw/docs/LORAS.md +131 -0
- dw/docs/MCP.md +517 -0
- dw/docs/PROMPT_WEIGHTING.md +78 -0
- dw/docs/QUANTIZATION.md +230 -0
- dw/docs/RECIPES_24GB.md +201 -0
- dw/docs/RELEASING.md +195 -0
- dw/docs/REMOTE.md +140 -0
- dw/docs/REPL_COMMANDS.md +121 -0
- dw/docs/REPL_WORKER_GUIDE.md +51 -0
- dw/docs/SECURITY.md +272 -0
- dw/docs/SECURITY_QUICKREF.md +112 -0
- dw/docs/SERVER.md +679 -0
- dw/docs/TASKS.md +1741 -0
- dw/docs/TESTING.md +71 -0
- dw/docs/WORKFLOW_GUIDE.md +2038 -0
- dw/docs/WORKSPACES.md +316 -0
- dw/download_watch.py +335 -0
- dw/elision.py +306 -0
- dw/events.py +275 -0
- dw/for_each.py +409 -0
- dw/host_memory.py +258 -0
- dw/host_memory_projection.py +230 -0
- dw/hub_cache.py +432 -0
- dw/introspection.py +1228 -0
- dw/kernel_availability.py +208 -0
- dw/locations.py +599 -0
- dw/log_setup.py +45 -0
- dw/loudness.py +82 -0
- dw/media_audio.py +217 -0
- dw/media_frames.py +367 -0
- dw/media_info.py +297 -0
- dw/pipeline_processors/chain.py +821 -0
- dw/pipeline_processors/config_objects.py +237 -0
- dw/pipeline_processors/pipeline.py +2297 -0
- dw/pipeline_processors/remote.py +46 -0
- dw/plan.py +920 -0
- dw/previous_results.py +411 -0
- dw/probe_paths.py +59 -0
- dw/prompt_schema.json +48 -0
- dw/prompt_weighting.py +378 -0
- dw/prompts.py +159 -0
- dw/realize.py +250 -0
- dw/reference_limits.py +215 -0
- dw/reference_names.py +125 -0
- dw/repl.py +338 -0
- dw/repl_commands.py +836 -0
- dw/repl_worker.py +159 -0
- dw/result.py +1720 -0
- dw/result_fps.py +82 -0
- dw/run.py +162 -0
- dw/runs.py +768 -0
- dw/scalar_result_validation.py +97 -0
- dw/schema.py +283 -0
- dw/security.py +1038 -0
- dw/select_validation.py +115 -0
- dw/serve.py +277 -0
- dw/server/__init__.py +2 -0
- dw/server/app.py +4586 -0
- dw/server/assess.py +132 -0
- dw/server/catalog_shape.py +487 -0
- dw/server/enhancers.py +129 -0
- dw/server/exports.py +480 -0
- dw/server/guides.py +257 -0
- dw/server/jobs.py +1561 -0
- dw/server/mcp_mount.py +95 -0
- dw/server/netinfo.py +124 -0
- dw/server/observed_cost.py +379 -0
- dw/server/sysinfo.py +71 -0
- dw/server/ui/assets/abap-08VXUWAP.js +1 -0
- dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
- dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
- dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
- dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
- dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
- dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
- dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
- dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
- dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
- dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
- dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
- dw/server/ui/assets/css-DIMkf-bt.js +3 -0
- dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
- dw/server/ui/assets/cssMode-CPznxfY8.js +1 -0
- dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
- dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
- dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
- dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
- dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
- dw/server/ui/assets/editor.api-CpWcotrd.js +847 -0
- dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
- dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
- dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
- dw/server/ui/assets/freemarker2-CXtRM8N4.js +3 -0
- dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
- dw/server/ui/assets/go-C-y9NEjX.js +1 -0
- dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
- dw/server/ui/assets/handlebars-N7x-6NMY.js +1 -0
- dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
- dw/server/ui/assets/html-PhsdjHSr.js +1 -0
- dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
- dw/server/ui/assets/htmlMode-Dgj0SEok.js +1 -0
- dw/server/ui/assets/index-3Vw6WAPW.css +1 -0
- dw/server/ui/assets/index-DgrYhQd9.js +43 -0
- dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
- dw/server/ui/assets/java-BEtHBSE6.js +1 -0
- dw/server/ui/assets/javascript-BJqN9Qhv.js +1 -0
- dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
- dw/server/ui/assets/jsonMode-DbM4SWSv.js +7 -0
- dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
- dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
- dw/server/ui/assets/less-B9JPFI3C.js +2 -0
- dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
- dw/server/ui/assets/liquid-BWr8lEc4.js +1 -0
- dw/server/ui/assets/lspLanguageFeatures-C1iGuDyZ.js +4 -0
- dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
- dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
- dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
- dw/server/ui/assets/mdx-DAdMi_0p.js +1 -0
- dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
- dw/server/ui/assets/monaco--ixms01u.css +1 -0
- dw/server/ui/assets/monaco-BGCeEqaw.js +56 -0
- dw/server/ui/assets/msdax-DauUninz.js +1 -0
- dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
- dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
- dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
- dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
- dw/server/ui/assets/perl-oz_6vUea.js +1 -0
- dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
- dw/server/ui/assets/php-nr791fC2.js +1 -0
- dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
- dw/server/ui/assets/postiats-43DmfD33.js +1 -0
- dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
- dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
- dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
- dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
- dw/server/ui/assets/python-Bcn70HdC.js +1 -0
- dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
- dw/server/ui/assets/r-BwWrilGY.js +1 -0
- dw/server/ui/assets/razor-D1HmNnby.js +1 -0
- dw/server/ui/assets/redis-ClamHrr6.js +1 -0
- dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
- dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
- dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
- dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
- dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
- dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
- dw/server/ui/assets/scheme-BeGwcela.js +1 -0
- dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
- dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
- dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
- dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
- dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
- dw/server/ui/assets/sql-NEE52Syq.js +1 -0
- dw/server/ui/assets/st-DbInun42.js +1 -0
- dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
- dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
- dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
- dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
- dw/server/ui/assets/tsMode-D6u0XmOW.js +11 -0
- dw/server/ui/assets/twig-De2hgUGE.js +1 -0
- dw/server/ui/assets/typescript-BU6v-LMV.js +1 -0
- dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
- dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
- dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
- dw/server/ui/assets/workers-Cn7cTUKr.js +1 -0
- dw/server/ui/assets/xml--0LP2Lwk.js +1 -0
- dw/server/ui/assets/yaml-mpBg9jnt.js +1 -0
- dw/server/ui/index.html +17 -0
- dw/server/updater.py +192 -0
- dw/settings.py +98 -0
- dw/shot_span_preflight.py +116 -0
- dw/shots.py +359 -0
- dw/slice_preflight.py +148 -0
- dw/step.py +187 -0
- dw/step_cache.py +442 -0
- dw/subfolders.py +107 -0
- dw/task_domains.py +307 -0
- dw/tasks/assess.py +826 -0
- dw/tasks/audio_transcription.py +88 -0
- dw/tasks/audio_utils.py +1862 -0
- dw/tasks/background_remover.py +43 -0
- dw/tasks/borders.py +113 -0
- dw/tasks/compose_text.py +74 -0
- dw/tasks/concat_videos.py +300 -0
- dw/tasks/depth_estimator.py +54 -0
- dw/tasks/diffusion_upscale.py +109 -0
- dw/tasks/dissolve_videos.py +342 -0
- dw/tasks/format_messages.py +24 -0
- dw/tasks/gather.py +173 -0
- dw/tasks/grade.py +97 -0
- dw/tasks/image_to_text.py +43 -0
- dw/tasks/image_utils.py +764 -0
- dw/tasks/interpolate_frames.py +252 -0
- dw/tasks/judge.py +68 -0
- dw/tasks/model_cache.py +55 -0
- dw/tasks/pair_audio.py +268 -0
- dw/tasks/qr_code.py +19 -0
- dw/tasks/restore_faces.py +175 -0
- dw/tasks/rife_model.py +192 -0
- dw/tasks/segment.py +121 -0
- dw/tasks/select.py +111 -0
- dw/tasks/speech_generation.py +228 -0
- dw/tasks/stabilize.py +129 -0
- dw/tasks/task.py +920 -0
- dw/tasks/tensor_image.py +57 -0
- dw/tasks/text_generation.py +169 -0
- dw/tasks/text_sections.py +80 -0
- dw/tasks/upscale.py +203 -0
- dw/tasks/video_utils.py +624 -0
- dw/tasks/zoe_depth.py +71 -0
- dw/teacache.py +381 -0
- dw/teacache_models.json +99 -0
- dw/test.py +29 -0
- dw/type_helpers.py +231 -0
- dw/validate.py +68 -0
- dw/variable_constraints.py +444 -0
- dw/variables.py +443 -0
- dw/video_extensions.py +141 -0
- dw/vram_estimate.py +116 -0
- dw/worker.py +764 -0
- dw/workflow.py +2007 -0
- dw/workflow_schema.json +1346 -0
- dw/workflow_sources.py +383 -0
- dw/workflows/h3_context_ir.json +57 -0
- dw/workflows/test.json +31 -0
- dw/workspace.py +730 -0
- dw_mcp/__init__.py +6 -0
- dw_mcp/__main__.py +133 -0
- dw_mcp/assets.py +336 -0
- dw_mcp/authoring.py +114 -0
- dw_mcp/catalog.py +360 -0
- dw_mcp/client.py +486 -0
- dw_mcp/diagnose.py +371 -0
- dw_mcp/exports.py +84 -0
- dw_mcp/guides.py +35 -0
- dw_mcp/media.py +638 -0
- dw_mcp/models.py +97 -0
- dw_mcp/prompts.py +104 -0
- dw_mcp/server.py +1343 -0
- dw_mcp/workspaces.py +212 -0
dw/for_each.py
ADDED
|
@@ -0,0 +1,409 @@
|
|
|
1
|
+
"""Expand a step's 'for_each' list into one ordinary step per entry.
|
|
2
|
+
|
|
3
|
+
A template that generates one shot per entry of a list is otherwise written
|
|
4
|
+
the long way - 'shot_1', 'shot_2', ... each a near-copy of the one before -
|
|
5
|
+
and a six-shot episode is a different file from a five-shot one. This pass
|
|
6
|
+
runs on the definition after variable substitution and before the run id is
|
|
7
|
+
computed (Workflow.run) and before the reference check (validation_errors),
|
|
8
|
+
and produces a definition with no 'for_each' in it: every member is an
|
|
9
|
+
ordinary step, so the step loop, the step cache, the manifest and the
|
|
10
|
+
reference checker never learn a new reference kind.
|
|
11
|
+
|
|
12
|
+
Inside a member:
|
|
13
|
+
- 'item:' is the whole entry; 'item:field' one field of an object
|
|
14
|
+
entry, spliced in whole whatever its type
|
|
15
|
+
- a reference to another for_each group over the SAME list resolves to
|
|
16
|
+
the member with the same key ('slice' inside 'shot@open' -> 'slice@open')
|
|
17
|
+
Outside (or inside, for any other group):
|
|
18
|
+
- 'gather:shot' is the list of every member's result, as explicit
|
|
19
|
+
'previous_result:shot@<key>' strings; inside a list it splices
|
|
20
|
+
- 'previous_result:shot' naming a group is an error that says to gather
|
|
21
|
+
|
|
22
|
+
Members are named '<group>@<key>' - the entry's own 'name' when it carries
|
|
23
|
+
one, else its index - and '@' is reserved in every step name. Names rather
|
|
24
|
+
than indexes because the step cache (dw/step_cache.py) keys on the step
|
|
25
|
+
name: inserting a shot in the middle of a list must not shift every later
|
|
26
|
+
member onto a different entry's cache line.
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
import copy
|
|
30
|
+
|
|
31
|
+
from .arguments import FROM_PREVIOUS_RESULT_KEY, PREVIOUS_RESULT_PREFIX
|
|
32
|
+
from .security import InvalidInputError, validate_variable_name
|
|
33
|
+
from .step_cache import reference_resolves_to
|
|
34
|
+
from .variables import argument_errors, set_variables
|
|
35
|
+
|
|
36
|
+
FOR_EACH_KEY = "for_each"
|
|
37
|
+
ITEM_PREFIX = "item:"
|
|
38
|
+
GATHER_PREFIX = "gather:"
|
|
39
|
+
MEMBER_SEPARATOR = "@"
|
|
40
|
+
# Each entry is a full generation. Stated against the step cache's bound
|
|
41
|
+
# (DEFAULT_MAX_ENTRIES = 128): a run whose expanded steps exceed the cache
|
|
42
|
+
# evicts its own earlier members, so this is kept well under it
|
|
43
|
+
MAX_FOR_EACH_ENTRIES = 32
|
|
44
|
+
# release_pipeline / release_models would drop the model after the first
|
|
45
|
+
# member and reload it for the second, so they are carried onto the last one
|
|
46
|
+
_LAST_MEMBER_ONLY = ("release_pipeline", "release_models")
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
class ForEachError(ValueError):
|
|
50
|
+
"""A for_each step that cannot be expanded, with the JSON path at fault."""
|
|
51
|
+
|
|
52
|
+
def __init__(self, path, message):
|
|
53
|
+
super().__init__(message)
|
|
54
|
+
self.path = path
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def member_name(group, key):
|
|
58
|
+
return f"{group}{MEMBER_SEPARATOR}{key}"
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def expand_for_each(definition, source_indices=None):
|
|
62
|
+
"""The definition with every 'for_each' step replaced by its members.
|
|
63
|
+
|
|
64
|
+
Returns a new structure; `definition` is left as it was passed in.
|
|
65
|
+
Raises ForEachError for anything that cannot be expanded.
|
|
66
|
+
|
|
67
|
+
`source_indices`, when a list is passed, has the index of the step each
|
|
68
|
+
expanded step was written as appended to it, so a later check can report
|
|
69
|
+
an error at a path in the file the author wrote rather than at an
|
|
70
|
+
expanded index that exists nowhere. A parallel list rather than a key on
|
|
71
|
+
the step: step_data is what the step cache keys on and what the schema
|
|
72
|
+
validates, and neither may learn a new field.
|
|
73
|
+
"""
|
|
74
|
+
steps = definition.get("steps") if isinstance(definition, dict) else None
|
|
75
|
+
if not isinstance(steps, list):
|
|
76
|
+
return definition
|
|
77
|
+
|
|
78
|
+
# group name -> {"keys": [...], "entries": [...]} for every group
|
|
79
|
+
# expanded so far, in step order, so a reference can only reach an
|
|
80
|
+
# earlier group - the same rule previous_result: has always had
|
|
81
|
+
groups = {}
|
|
82
|
+
expanded = []
|
|
83
|
+
for index, step in enumerate(steps):
|
|
84
|
+
if not isinstance(step, dict):
|
|
85
|
+
expanded.append(copy.deepcopy(step))
|
|
86
|
+
_record(source_indices, index)
|
|
87
|
+
continue
|
|
88
|
+
path = ("steps", index)
|
|
89
|
+
name = step.get("name")
|
|
90
|
+
if isinstance(name, str) and MEMBER_SEPARATOR in name:
|
|
91
|
+
raise ForEachError(
|
|
92
|
+
render_path(path + ("name",)),
|
|
93
|
+
f"Step name '{name}' contains '{MEMBER_SEPARATOR}', which is "
|
|
94
|
+
f"reserved for the members of a for_each step",
|
|
95
|
+
)
|
|
96
|
+
if FOR_EACH_KEY not in step:
|
|
97
|
+
expanded.append(_rewrite(step, path, groups, member=None))
|
|
98
|
+
_record(source_indices, index)
|
|
99
|
+
continue
|
|
100
|
+
|
|
101
|
+
entries = step[FOR_EACH_KEY]
|
|
102
|
+
keys = _entry_keys(entries, path + (FOR_EACH_KEY,))
|
|
103
|
+
template = {k: v for k, v in step.items() if k != FOR_EACH_KEY}
|
|
104
|
+
last = len(keys) - 1
|
|
105
|
+
for position, (key, entry) in enumerate(zip(keys, entries)):
|
|
106
|
+
member = {
|
|
107
|
+
"group": name,
|
|
108
|
+
"key": key,
|
|
109
|
+
"entry": entry,
|
|
110
|
+
"entries": entries,
|
|
111
|
+
"index": position,
|
|
112
|
+
}
|
|
113
|
+
expanded_step = _rewrite(template, path, groups, member)
|
|
114
|
+
expanded_step["name"] = member_name(name, key)
|
|
115
|
+
if position != last:
|
|
116
|
+
for flag in _LAST_MEMBER_ONLY:
|
|
117
|
+
expanded_step.pop(flag, None)
|
|
118
|
+
expanded.append(expanded_step)
|
|
119
|
+
_record(source_indices, index)
|
|
120
|
+
groups[name] = {"keys": keys, "entries": entries}
|
|
121
|
+
|
|
122
|
+
result = {k: v for k, v in definition.items() if k != "steps"}
|
|
123
|
+
result["steps"] = expanded
|
|
124
|
+
return result
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
def _record(source_indices, index):
|
|
128
|
+
if source_indices is not None:
|
|
129
|
+
source_indices.append(index)
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
def _entry_keys(entries, path):
|
|
133
|
+
"""The key of every entry - its 'name' when it is an object carrying
|
|
134
|
+
one, else its index - validated and unique."""
|
|
135
|
+
if not isinstance(entries, list):
|
|
136
|
+
if isinstance(entries, str) and entries.startswith("variable:"):
|
|
137
|
+
hint = f" - '{entries}' was not substituted; is the variable declared?"
|
|
138
|
+
else:
|
|
139
|
+
hint = ""
|
|
140
|
+
raise ForEachError(
|
|
141
|
+
render_path(path),
|
|
142
|
+
f"for_each must be a list, got {type(entries).__name__}{hint}",
|
|
143
|
+
)
|
|
144
|
+
if not entries:
|
|
145
|
+
raise ForEachError(
|
|
146
|
+
render_path(path),
|
|
147
|
+
"for_each over an empty list would run no steps - a workflow that "
|
|
148
|
+
"generates nothing is never what was asked for",
|
|
149
|
+
)
|
|
150
|
+
if len(entries) > MAX_FOR_EACH_ENTRIES:
|
|
151
|
+
raise ForEachError(
|
|
152
|
+
render_path(path),
|
|
153
|
+
f"for_each has {len(entries)} entries; the limit is {MAX_FOR_EACH_ENTRIES}",
|
|
154
|
+
)
|
|
155
|
+
keys = []
|
|
156
|
+
for index, entry in enumerate(entries):
|
|
157
|
+
key = str(index)
|
|
158
|
+
if isinstance(entry, dict) and "name" in entry:
|
|
159
|
+
key = entry["name"]
|
|
160
|
+
key_path = render_path(path + (index, "name"))
|
|
161
|
+
if not isinstance(key, str):
|
|
162
|
+
raise ForEachError(key_path, "An entry's name must be a string")
|
|
163
|
+
try:
|
|
164
|
+
validate_variable_name(key)
|
|
165
|
+
except InvalidInputError as e:
|
|
166
|
+
raise ForEachError(key_path, f"Invalid entry name '{key}': {e}") from e
|
|
167
|
+
if key in keys:
|
|
168
|
+
raise ForEachError(key_path, f"Duplicate entry name '{key}'")
|
|
169
|
+
keys.append(key)
|
|
170
|
+
return keys
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
def _rewrite(value, path, groups, member):
|
|
174
|
+
"""Rebuild `value` with item:, gather: and group references resolved.
|
|
175
|
+
|
|
176
|
+
`member` is None outside a for_each step; inside one it carries the
|
|
177
|
+
group, key and entry of the member being built.
|
|
178
|
+
"""
|
|
179
|
+
if isinstance(value, dict):
|
|
180
|
+
rebuilt = {}
|
|
181
|
+
for key, item in value.items():
|
|
182
|
+
if key == FROM_PREVIOUS_RESULT_KEY and isinstance(item, str):
|
|
183
|
+
rebuilt[key] = _rewrite_reference(item, path + (key,), groups, member)
|
|
184
|
+
else:
|
|
185
|
+
rebuilt[key] = _rewrite(item, path + (key,), groups, member)
|
|
186
|
+
return rebuilt
|
|
187
|
+
if isinstance(value, list):
|
|
188
|
+
rebuilt = []
|
|
189
|
+
for index, item in enumerate(value):
|
|
190
|
+
if isinstance(item, str) and item.startswith(GATHER_PREFIX):
|
|
191
|
+
# A gather inside a list splices into it
|
|
192
|
+
rebuilt.extend(_gather(item, path + (index,), groups))
|
|
193
|
+
else:
|
|
194
|
+
rebuilt.append(_rewrite(item, path + (index,), groups, member))
|
|
195
|
+
return rebuilt
|
|
196
|
+
if isinstance(value, str):
|
|
197
|
+
if value.startswith(GATHER_PREFIX):
|
|
198
|
+
return _gather(value, path, groups)
|
|
199
|
+
if value.startswith(ITEM_PREFIX):
|
|
200
|
+
return _item(value, path, member)
|
|
201
|
+
if value.startswith(PREVIOUS_RESULT_PREFIX):
|
|
202
|
+
reference = value[len(PREVIOUS_RESULT_PREFIX) :]
|
|
203
|
+
return PREVIOUS_RESULT_PREFIX + _rewrite_reference(
|
|
204
|
+
reference, path, groups, member
|
|
205
|
+
)
|
|
206
|
+
return value
|
|
207
|
+
# A leaf is only copied where the copy is needed: inside a member, where
|
|
208
|
+
# the same template value is about to appear in every one of them.
|
|
209
|
+
# Outside, the leaf is handed back as it is - this pass runs on every run
|
|
210
|
+
# of every workflow, after realize_args has turned 'asset:' arguments
|
|
211
|
+
# into loaded images and decoded frame lists, and copying all of that
|
|
212
|
+
# would multiply the media a run holds. The 'input untouched' contract
|
|
213
|
+
# still holds because nothing here ever mutates a leaf
|
|
214
|
+
return _copy_leaf(value) if member is not None else value
|
|
215
|
+
|
|
216
|
+
|
|
217
|
+
def _copy_leaf(value):
|
|
218
|
+
"""A copy of a leaf, or the leaf itself when it cannot be copied.
|
|
219
|
+
|
|
220
|
+
An open handle or a live model object reaching a member is not a reason
|
|
221
|
+
to fail a run - the step cache makes the same choice for a realized
|
|
222
|
+
argument it cannot deep-copy (dw/workflow.py).
|
|
223
|
+
"""
|
|
224
|
+
try:
|
|
225
|
+
return copy.deepcopy(value)
|
|
226
|
+
except Exception:
|
|
227
|
+
return value
|
|
228
|
+
|
|
229
|
+
|
|
230
|
+
def _item(value, path, member):
|
|
231
|
+
if member is None:
|
|
232
|
+
raise ForEachError(
|
|
233
|
+
render_path(path), f"'{value}' is only meaningful inside a for_each step"
|
|
234
|
+
)
|
|
235
|
+
field = value[len(ITEM_PREFIX) :]
|
|
236
|
+
entry = member["entry"]
|
|
237
|
+
if field == "":
|
|
238
|
+
return _copy_leaf(entry)
|
|
239
|
+
if not isinstance(entry, dict):
|
|
240
|
+
raise ForEachError(
|
|
241
|
+
render_path(path),
|
|
242
|
+
f"'{value}' asks for a field of entry '{member['key']}' of "
|
|
243
|
+
f"for_each step '{member['group']}', which is not an object",
|
|
244
|
+
)
|
|
245
|
+
if field not in entry:
|
|
246
|
+
raise ForEachError(
|
|
247
|
+
render_path(path),
|
|
248
|
+
f"'{value}' names no field of entry '{member['key']}' of for_each "
|
|
249
|
+
f"step '{member['group']}'; it has: {sorted(entry)}",
|
|
250
|
+
)
|
|
251
|
+
return _copy_leaf(entry[field])
|
|
252
|
+
|
|
253
|
+
|
|
254
|
+
def _gather(value, path, groups):
|
|
255
|
+
group = value[len(GATHER_PREFIX) :]
|
|
256
|
+
if group not in groups:
|
|
257
|
+
raise ForEachError(
|
|
258
|
+
render_path(path),
|
|
259
|
+
f"'{value}' names no earlier for_each step. "
|
|
260
|
+
f"for_each steps available here: {sorted(groups)}",
|
|
261
|
+
)
|
|
262
|
+
return [
|
|
263
|
+
PREVIOUS_RESULT_PREFIX + member_name(group, key)
|
|
264
|
+
for key in groups[group]["keys"]
|
|
265
|
+
]
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
def _rewrite_reference(reference, path, groups, member):
|
|
269
|
+
"""A previous_result reference (without its prefix) as the expanded
|
|
270
|
+
definition spells it: unchanged unless it names a for_each group."""
|
|
271
|
+
if reference.startswith("variable:"):
|
|
272
|
+
return reference
|
|
273
|
+
group = next((g for g in groups if reference_resolves_to(reference, g)), None)
|
|
274
|
+
if group is None:
|
|
275
|
+
if member is not None and reference_resolves_to(reference, member["group"]):
|
|
276
|
+
raise ForEachError(
|
|
277
|
+
render_path(path),
|
|
278
|
+
f"'{reference}' names its own for_each step '{member['group']}'",
|
|
279
|
+
)
|
|
280
|
+
return reference
|
|
281
|
+
if member is not None and member["entries"] == groups[group]["entries"]:
|
|
282
|
+
# Same list: shot@open reads slice@open
|
|
283
|
+
return member_name(group, member["key"]) + reference[len(group) :]
|
|
284
|
+
where = (
|
|
285
|
+
f"from inside for_each step '{member['group']}', which runs over a different list"
|
|
286
|
+
if member is not None
|
|
287
|
+
else "from outside a for_each step"
|
|
288
|
+
)
|
|
289
|
+
raise ForEachError(
|
|
290
|
+
render_path(path),
|
|
291
|
+
f"'{reference}' names the for_each step '{group}' {where}. Use "
|
|
292
|
+
f"'{GATHER_PREFIX}{group}' for every member's result, or a reference "
|
|
293
|
+
f"from a for_each step over the same list for the same-keyed member",
|
|
294
|
+
)
|
|
295
|
+
|
|
296
|
+
|
|
297
|
+
def list_fields(definition):
|
|
298
|
+
"""What an entry of each list-driven variable has to carry, read off
|
|
299
|
+
the definition: for every step whose for_each is 'variable:<name>',
|
|
300
|
+
the fields its 'item:<field>' references name.
|
|
301
|
+
|
|
302
|
+
Returns {<variable>: {"fields": [...] or None, "steps": [...]}} -
|
|
303
|
+
fields sorted with 'name' first, or None when a step splices the
|
|
304
|
+
whole entry with a bare 'item:' (the entries are values, not
|
|
305
|
+
objects). A literal for_each list is not an argument and is skipped.
|
|
306
|
+
Reads the raw definition, no substitution, so the catalog and the
|
|
307
|
+
validator derive the same answer from the file as written.
|
|
308
|
+
"""
|
|
309
|
+
steps = definition.get("steps") if isinstance(definition, dict) else None
|
|
310
|
+
if not isinstance(steps, list):
|
|
311
|
+
return {}
|
|
312
|
+
found = {}
|
|
313
|
+
for step in steps:
|
|
314
|
+
if not isinstance(step, dict):
|
|
315
|
+
continue
|
|
316
|
+
target = step.get(FOR_EACH_KEY)
|
|
317
|
+
if not (isinstance(target, str) and target.startswith("variable:")):
|
|
318
|
+
continue
|
|
319
|
+
variable = target.removeprefix("variable:")
|
|
320
|
+
entry = found.setdefault(variable, {"fields": set(), "steps": []})
|
|
321
|
+
entry["steps"].append(step.get("name"))
|
|
322
|
+
for value in _strings(step):
|
|
323
|
+
if not value.startswith(ITEM_PREFIX):
|
|
324
|
+
continue
|
|
325
|
+
field = value[len(ITEM_PREFIX) :]
|
|
326
|
+
if field == "":
|
|
327
|
+
entry["fields"] = None
|
|
328
|
+
elif entry["fields"] is not None:
|
|
329
|
+
entry["fields"].add(field)
|
|
330
|
+
return {
|
|
331
|
+
variable: {
|
|
332
|
+
"fields": (
|
|
333
|
+
None
|
|
334
|
+
if entry["fields"] is None
|
|
335
|
+
else ["name"] + sorted(entry["fields"] - {"name"})
|
|
336
|
+
),
|
|
337
|
+
"steps": entry["steps"],
|
|
338
|
+
}
|
|
339
|
+
for variable, entry in found.items()
|
|
340
|
+
}
|
|
341
|
+
|
|
342
|
+
|
|
343
|
+
def entry_field_warnings(definition, arguments=None):
|
|
344
|
+
"""Every entry key of a list-driven variable that no step reads.
|
|
345
|
+
|
|
346
|
+
A caller who writes 'num_frame' for 'num_frames' gets the template's
|
|
347
|
+
value for the field they meant to set, in silence; this names the
|
|
348
|
+
key, at the entry it sits in, with the fields the list takes. A
|
|
349
|
+
warning rather than an error: an entry may carry a note on purpose.
|
|
350
|
+
Good `arguments` are folded in first, and a list the caller supplied
|
|
351
|
+
is reported under 'arguments.', where they wrote it.
|
|
352
|
+
"""
|
|
353
|
+
if not isinstance(definition, dict):
|
|
354
|
+
return []
|
|
355
|
+
variables = definition.get("variables")
|
|
356
|
+
if not isinstance(variables, dict):
|
|
357
|
+
return []
|
|
358
|
+
variables = copy.deepcopy(variables)
|
|
359
|
+
supplied = set()
|
|
360
|
+
if arguments and not argument_errors(definition, arguments):
|
|
361
|
+
set_variables(arguments, variables)
|
|
362
|
+
supplied = set(arguments)
|
|
363
|
+
warnings = []
|
|
364
|
+
for variable, spec in list_fields(definition).items():
|
|
365
|
+
fields = spec["fields"]
|
|
366
|
+
entries = variables.get(variable)
|
|
367
|
+
if fields is None or not isinstance(entries, list):
|
|
368
|
+
continue
|
|
369
|
+
where = "arguments" if variable in supplied else "variables"
|
|
370
|
+
for index, entry in enumerate(entries):
|
|
371
|
+
if not isinstance(entry, dict):
|
|
372
|
+
continue
|
|
373
|
+
unknown = sorted(set(entry) - set(fields))
|
|
374
|
+
if not unknown:
|
|
375
|
+
continue
|
|
376
|
+
label = repr(entry["name"]) if isinstance(entry.get("name"), str) else index
|
|
377
|
+
warnings.append(
|
|
378
|
+
f"{where}.{variable}[{index}]: entry {label} carries "
|
|
379
|
+
f"{', '.join(repr(k) for k in unknown)}, which no step reads; "
|
|
380
|
+
f"entries of '{variable}' take: {', '.join(fields)}"
|
|
381
|
+
)
|
|
382
|
+
return warnings
|
|
383
|
+
|
|
384
|
+
|
|
385
|
+
def _strings(value):
|
|
386
|
+
"""Every string anywhere inside a JSON value, except the for_each key
|
|
387
|
+
itself."""
|
|
388
|
+
if isinstance(value, str):
|
|
389
|
+
yield value
|
|
390
|
+
elif isinstance(value, list):
|
|
391
|
+
for item in value:
|
|
392
|
+
yield from _strings(item)
|
|
393
|
+
elif isinstance(value, dict):
|
|
394
|
+
for key, item in value.items():
|
|
395
|
+
if key != FOR_EACH_KEY:
|
|
396
|
+
yield from _strings(item)
|
|
397
|
+
|
|
398
|
+
|
|
399
|
+
def render_path(path):
|
|
400
|
+
"""'steps[3].task.arguments.videos[1]' - the same shape schema errors use."""
|
|
401
|
+
rendered = ""
|
|
402
|
+
for part in path:
|
|
403
|
+
if isinstance(part, int):
|
|
404
|
+
rendered += f"[{part}]"
|
|
405
|
+
elif rendered:
|
|
406
|
+
rendered += f".{part}"
|
|
407
|
+
else:
|
|
408
|
+
rendered = str(part)
|
|
409
|
+
return rendered
|
dw/host_memory.py
ADDED
|
@@ -0,0 +1,258 @@
|
|
|
1
|
+
"""What the worker process and its machine are using in host RAM.
|
|
2
|
+
|
|
3
|
+
`device_memory_stats` answers for the accelerator, which on this project is
|
|
4
|
+
often the least informative half of the question: the templates here keep
|
|
5
|
+
weights in host memory by design (`offload: "sequential"`, `group_offload`),
|
|
6
|
+
so a card can sit near-empty through a generation whose weights are very
|
|
7
|
+
much resident somewhere. This module is the other half.
|
|
8
|
+
|
|
9
|
+
Three methods, in order, none of them required: psutil when it happens to be
|
|
10
|
+
installed (a transitive dependency here, not a declared one), Linux's
|
|
11
|
+
`/proc` otherwise, and `resource.getrusage` for the process's peak, which is
|
|
12
|
+
available on every POSIX platform. Nothing raises - a reading that cannot be
|
|
13
|
+
taken is reported as None rather than failing the call that asked for it.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
import logging
|
|
17
|
+
import os
|
|
18
|
+
|
|
19
|
+
logger = logging.getLogger("dw")
|
|
20
|
+
|
|
21
|
+
__all__ = [
|
|
22
|
+
"host_memory_stats",
|
|
23
|
+
"host_memory_fields",
|
|
24
|
+
"trim_host_memory",
|
|
25
|
+
"release_host_caches",
|
|
26
|
+
"pinned_host_memory_fields",
|
|
27
|
+
]
|
|
28
|
+
|
|
29
|
+
_MB = 1024.0 * 1024.0
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def host_memory_stats():
|
|
33
|
+
"""Snapshot of host memory for this process and the machine it is on.
|
|
34
|
+
|
|
35
|
+
Returns:
|
|
36
|
+
dict with keys, any of which may be None when the platform cannot
|
|
37
|
+
answer:
|
|
38
|
+
rss_mb (float or None): resident set size of *this* process -
|
|
39
|
+
in the worker, the weights it is holding
|
|
40
|
+
peak_rss_mb (float or None): the high-water mark of the above,
|
|
41
|
+
which is what says whether a run that has finished released
|
|
42
|
+
what it took
|
|
43
|
+
total_mb (float or None): the machine's physical memory
|
|
44
|
+
available_mb (float or None): what can be handed out without
|
|
45
|
+
swapping - MemAvailable on Linux, not free memory, since
|
|
46
|
+
page cache is reclaimable
|
|
47
|
+
"""
|
|
48
|
+
stats = {
|
|
49
|
+
"rss_mb": None,
|
|
50
|
+
"peak_rss_mb": _peak_rss_mb(),
|
|
51
|
+
"total_mb": None,
|
|
52
|
+
"available_mb": None,
|
|
53
|
+
}
|
|
54
|
+
|
|
55
|
+
for method in (_psutil_stats, _proc_stats):
|
|
56
|
+
try:
|
|
57
|
+
reading = method()
|
|
58
|
+
except Exception as e: # a memory reading is never worth an exception
|
|
59
|
+
logger.debug(f"Host memory via {method.__name__} failed: {e}")
|
|
60
|
+
continue
|
|
61
|
+
for key, value in reading.items():
|
|
62
|
+
if stats.get(key) is None and value is not None:
|
|
63
|
+
stats[key] = value
|
|
64
|
+
if all(stats[key] is not None for key in ("rss_mb", "total_mb")):
|
|
65
|
+
break
|
|
66
|
+
|
|
67
|
+
return _hold_the_high_water_mark(stats)
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def _hold_the_high_water_mark(stats):
|
|
71
|
+
"""Keep peak_rss_mb >= rss_mb, which is what a high-water mark means.
|
|
72
|
+
|
|
73
|
+
The two readings come from different places - getrusage's ru_maxrss,
|
|
74
|
+
quantized to whole pages and taken first, against psutil's rss taken a
|
|
75
|
+
moment later - so a process that has never peaked meaningfully above its
|
|
76
|
+
current size reports them within a megabyte of each other in either
|
|
77
|
+
order. `peak - rss` is the whole point of the pair (what a run took and
|
|
78
|
+
did not give back), and a small negative there reads as "these fields
|
|
79
|
+
are not comparable" rather than "nothing leaked" (#83).
|
|
80
|
+
"""
|
|
81
|
+
peak, rss = stats["peak_rss_mb"], stats["rss_mb"]
|
|
82
|
+
if peak is not None and rss is not None and peak < rss:
|
|
83
|
+
stats["peak_rss_mb"] = rss
|
|
84
|
+
return stats
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def _psutil_stats():
|
|
88
|
+
import psutil
|
|
89
|
+
|
|
90
|
+
virtual = psutil.virtual_memory()
|
|
91
|
+
return {
|
|
92
|
+
"rss_mb": psutil.Process().memory_info().rss / _MB,
|
|
93
|
+
"total_mb": virtual.total / _MB,
|
|
94
|
+
"available_mb": virtual.available / _MB,
|
|
95
|
+
}
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def _proc_stats():
|
|
99
|
+
"""Linux without psutil: /proc/self/statm for this process, /proc/meminfo
|
|
100
|
+
for the machine. Both absent elsewhere, which is what the None default
|
|
101
|
+
covers."""
|
|
102
|
+
reading = {"rss_mb": None, "total_mb": None, "available_mb": None}
|
|
103
|
+
|
|
104
|
+
try:
|
|
105
|
+
with open("/proc/self/statm", "r") as f:
|
|
106
|
+
pages = int(f.read().split()[1])
|
|
107
|
+
reading["rss_mb"] = pages * os.sysconf("SC_PAGE_SIZE") / _MB
|
|
108
|
+
except (OSError, ValueError, IndexError, AttributeError):
|
|
109
|
+
pass
|
|
110
|
+
|
|
111
|
+
wanted = {"MemTotal:": "total_mb", "MemAvailable:": "available_mb"}
|
|
112
|
+
try:
|
|
113
|
+
with open("/proc/meminfo", "r") as f:
|
|
114
|
+
for line in f:
|
|
115
|
+
parts = line.split()
|
|
116
|
+
if len(parts) >= 2 and parts[0] in wanted:
|
|
117
|
+
# meminfo is in kB
|
|
118
|
+
reading[wanted[parts[0]]] = int(parts[1]) / 1024.0
|
|
119
|
+
except (OSError, ValueError):
|
|
120
|
+
pass
|
|
121
|
+
|
|
122
|
+
return reading
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def _peak_rss_mb():
|
|
126
|
+
"""getrusage's high-water mark, in kB on Linux and bytes on macOS - the
|
|
127
|
+
one place in this module where the unit depends on the platform."""
|
|
128
|
+
try:
|
|
129
|
+
import resource
|
|
130
|
+
import sys
|
|
131
|
+
|
|
132
|
+
peak = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
|
|
133
|
+
except Exception:
|
|
134
|
+
return None
|
|
135
|
+
if not peak:
|
|
136
|
+
return None
|
|
137
|
+
return peak / _MB if sys.platform == "darwin" else peak / 1024.0
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
# The memory payload's own names for the above, beside its gpu_* keys
|
|
141
|
+
FIELD_NAMES = {
|
|
142
|
+
"rss_mb": "host_memory_rss_mb",
|
|
143
|
+
"peak_rss_mb": "host_memory_peak_rss_mb",
|
|
144
|
+
"total_mb": "host_memory_total_mb",
|
|
145
|
+
"available_mb": "host_memory_available_mb",
|
|
146
|
+
}
|
|
147
|
+
|
|
148
|
+
|
|
149
|
+
def host_memory_fields():
|
|
150
|
+
"""host_memory_stats under the names the memory payload reports, with
|
|
151
|
+
the readings this platform cannot take left out rather than sent as
|
|
152
|
+
null - a key that is absent says "not measurable here", where a null
|
|
153
|
+
would read as "measured, and nothing"."""
|
|
154
|
+
stats = host_memory_stats()
|
|
155
|
+
return {
|
|
156
|
+
FIELD_NAMES[key]: value for key, value in stats.items() if value is not None
|
|
157
|
+
}
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
def trim_host_memory():
|
|
161
|
+
"""Hand memory the process has already freed back to the operating
|
|
162
|
+
system, and report how much that was in MB (0.0 when the platform has no
|
|
163
|
+
way to ask).
|
|
164
|
+
|
|
165
|
+
Dropping the last reference to a model frees it inside the process, not
|
|
166
|
+
back to the kernel: glibc keeps the arenas the weights were read into and
|
|
167
|
+
hands them out again to *this* process. That is normally invisible and
|
|
168
|
+
correct - until a template needs 96% of host RAM to run at all, at which
|
|
169
|
+
point the several GB the previous job's arenas are sitting on is the
|
|
170
|
+
difference between a run and a SIGKILL five minutes in (#98). The
|
|
171
|
+
templates here load tens of GB of weights through host memory, so the
|
|
172
|
+
arenas in question are large ones and the fragmentation that keeps
|
|
173
|
+
`malloc_trim` from returning them is the exception rather than the rule.
|
|
174
|
+
|
|
175
|
+
Linux/glibc only: `malloc_trim` is a GNU extension. Everywhere else -
|
|
176
|
+
macOS included - this is a no-op that reports 0.0, because there is
|
|
177
|
+
nothing to ask and a fabricated number would be worse than none.
|
|
178
|
+
"""
|
|
179
|
+
import ctypes
|
|
180
|
+
import ctypes.util
|
|
181
|
+
import sys
|
|
182
|
+
|
|
183
|
+
if not sys.platform.startswith("linux"):
|
|
184
|
+
return 0.0
|
|
185
|
+
before = host_memory_stats()["rss_mb"]
|
|
186
|
+
try:
|
|
187
|
+
libc = ctypes.CDLL(ctypes.util.find_library("c") or "libc.so.6", use_errno=True)
|
|
188
|
+
libc.malloc_trim(ctypes.c_size_t(0))
|
|
189
|
+
except (OSError, AttributeError) as e:
|
|
190
|
+
# musl and friends have no malloc_trim; not having one is not an error
|
|
191
|
+
logger.debug(f"malloc_trim unavailable: {e}")
|
|
192
|
+
return 0.0
|
|
193
|
+
after = host_memory_stats()["rss_mb"]
|
|
194
|
+
if before is None or after is None:
|
|
195
|
+
return 0.0
|
|
196
|
+
return max(0.0, before - after)
|
|
197
|
+
|
|
198
|
+
|
|
199
|
+
def pinned_host_memory_fields():
|
|
200
|
+
"""What CUDA's pinned-host allocator is holding, in MB, or {} where
|
|
201
|
+
there is none to report.
|
|
202
|
+
|
|
203
|
+
Group offloading with `use_stream` stages a component's weights through
|
|
204
|
+
*pinned* host memory, which torch caches per process exactly as it
|
|
205
|
+
caches device memory: freeing the tensors returns the blocks to that
|
|
206
|
+
cache, not to the OS, so they stay in this process's RSS and count
|
|
207
|
+
against the next job's host budget. It is invisible in every figure the
|
|
208
|
+
memory payload carried before - `rss_mb` includes it without saying so
|
|
209
|
+
and `gpu_memory_*` does not see it at all (#98).
|
|
210
|
+
"""
|
|
211
|
+
try:
|
|
212
|
+
import torch
|
|
213
|
+
|
|
214
|
+
stats = torch.cuda.host_memory_stats()
|
|
215
|
+
except Exception:
|
|
216
|
+
return {}
|
|
217
|
+
fields = {}
|
|
218
|
+
for key, name in (
|
|
219
|
+
("allocated_bytes.all.current", "host_pinned_allocated_mb"),
|
|
220
|
+
("reserved_bytes.all.current", "host_pinned_reserved_mb"),
|
|
221
|
+
):
|
|
222
|
+
value = stats.get(key)
|
|
223
|
+
if value is not None:
|
|
224
|
+
fields[name] = value / _MB
|
|
225
|
+
return fields
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
def release_host_caches():
|
|
229
|
+
"""Give back host memory this process is holding but no longer using,
|
|
230
|
+
and report what came back in MB.
|
|
231
|
+
|
|
232
|
+
Two caches, neither of which `gc.collect()` touches:
|
|
233
|
+
|
|
234
|
+
- torch's pinned-host allocator, where group offloading's staging
|
|
235
|
+
buffers live. `_host_emptyCache` frees the blocks nothing is using;
|
|
236
|
+
blocks a still-loaded pipeline is staging through are in use and are
|
|
237
|
+
not touched, so this is safe to call with models resident.
|
|
238
|
+
- glibc's heap arenas, via `malloc_trim`. Freeing a large block inside
|
|
239
|
+
the process does not hand its pages back to the kernel.
|
|
240
|
+
|
|
241
|
+
Together they are why a worker that has released every model still sat
|
|
242
|
+
on 14.5 GB, which is the difference between the next job running and
|
|
243
|
+
being OOM-killed five minutes in on a template that needs 96% of host
|
|
244
|
+
RAM (#98).
|
|
245
|
+
"""
|
|
246
|
+
before = host_memory_stats()["rss_mb"]
|
|
247
|
+
try:
|
|
248
|
+
import torch
|
|
249
|
+
|
|
250
|
+
if hasattr(torch._C, "_host_emptyCache"):
|
|
251
|
+
torch._C._host_emptyCache()
|
|
252
|
+
except Exception as e: # a cleanup is never worth failing the run for
|
|
253
|
+
logger.debug(f"Could not empty the pinned host cache: {e}")
|
|
254
|
+
trim_host_memory()
|
|
255
|
+
after = host_memory_stats()["rss_mb"]
|
|
256
|
+
if before is None or after is None:
|
|
257
|
+
return 0.0
|
|
258
|
+
return max(0.0, before - after)
|