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.
Files changed (260) hide show
  1. diffusers_workflow-0.4.0.dist-info/METADATA +318 -0
  2. diffusers_workflow-0.4.0.dist-info/RECORD +260 -0
  3. diffusers_workflow-0.4.0.dist-info/WHEEL +5 -0
  4. diffusers_workflow-0.4.0.dist-info/entry_points.txt +7 -0
  5. diffusers_workflow-0.4.0.dist-info/licenses/LICENSE +201 -0
  6. diffusers_workflow-0.4.0.dist-info/top_level.txt +2 -0
  7. dw/__init__.py +440 -0
  8. dw/adapter_compatibility.py +226 -0
  9. dw/arguments.py +1231 -0
  10. dw/assessment_rules.py +159 -0
  11. dw/assets.py +130 -0
  12. dw/cache_blocks.json +16 -0
  13. dw/cache_blocks.py +146 -0
  14. dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
  15. dw/content_types.py +150 -0
  16. dw/dissolve_frame_errors.py +121 -0
  17. dw/docs/ACCELERATION.md +352 -0
  18. dw/docs/AGENT_LOOP.md +95 -0
  19. dw/docs/DEPENDENCIES.md +91 -0
  20. dw/docs/IP_ADAPTER.md +109 -0
  21. dw/docs/LORAS.md +131 -0
  22. dw/docs/MCP.md +517 -0
  23. dw/docs/PROMPT_WEIGHTING.md +78 -0
  24. dw/docs/QUANTIZATION.md +230 -0
  25. dw/docs/RECIPES_24GB.md +201 -0
  26. dw/docs/RELEASING.md +195 -0
  27. dw/docs/REMOTE.md +140 -0
  28. dw/docs/REPL_COMMANDS.md +121 -0
  29. dw/docs/REPL_WORKER_GUIDE.md +51 -0
  30. dw/docs/SECURITY.md +272 -0
  31. dw/docs/SECURITY_QUICKREF.md +112 -0
  32. dw/docs/SERVER.md +679 -0
  33. dw/docs/TASKS.md +1741 -0
  34. dw/docs/TESTING.md +71 -0
  35. dw/docs/WORKFLOW_GUIDE.md +2038 -0
  36. dw/docs/WORKSPACES.md +316 -0
  37. dw/download_watch.py +335 -0
  38. dw/elision.py +306 -0
  39. dw/events.py +275 -0
  40. dw/for_each.py +409 -0
  41. dw/host_memory.py +258 -0
  42. dw/host_memory_projection.py +230 -0
  43. dw/hub_cache.py +432 -0
  44. dw/introspection.py +1228 -0
  45. dw/kernel_availability.py +208 -0
  46. dw/locations.py +599 -0
  47. dw/log_setup.py +45 -0
  48. dw/loudness.py +82 -0
  49. dw/media_audio.py +217 -0
  50. dw/media_frames.py +367 -0
  51. dw/media_info.py +297 -0
  52. dw/pipeline_processors/chain.py +821 -0
  53. dw/pipeline_processors/config_objects.py +237 -0
  54. dw/pipeline_processors/pipeline.py +2297 -0
  55. dw/pipeline_processors/remote.py +46 -0
  56. dw/plan.py +920 -0
  57. dw/previous_results.py +411 -0
  58. dw/probe_paths.py +59 -0
  59. dw/prompt_schema.json +48 -0
  60. dw/prompt_weighting.py +378 -0
  61. dw/prompts.py +159 -0
  62. dw/realize.py +250 -0
  63. dw/reference_limits.py +215 -0
  64. dw/reference_names.py +125 -0
  65. dw/repl.py +338 -0
  66. dw/repl_commands.py +836 -0
  67. dw/repl_worker.py +159 -0
  68. dw/result.py +1720 -0
  69. dw/result_fps.py +82 -0
  70. dw/run.py +162 -0
  71. dw/runs.py +768 -0
  72. dw/scalar_result_validation.py +97 -0
  73. dw/schema.py +283 -0
  74. dw/security.py +1038 -0
  75. dw/select_validation.py +115 -0
  76. dw/serve.py +277 -0
  77. dw/server/__init__.py +2 -0
  78. dw/server/app.py +4586 -0
  79. dw/server/assess.py +132 -0
  80. dw/server/catalog_shape.py +487 -0
  81. dw/server/enhancers.py +129 -0
  82. dw/server/exports.py +480 -0
  83. dw/server/guides.py +257 -0
  84. dw/server/jobs.py +1561 -0
  85. dw/server/mcp_mount.py +95 -0
  86. dw/server/netinfo.py +124 -0
  87. dw/server/observed_cost.py +379 -0
  88. dw/server/sysinfo.py +71 -0
  89. dw/server/ui/assets/abap-08VXUWAP.js +1 -0
  90. dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
  91. dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
  92. dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
  93. dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
  94. dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
  95. dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
  96. dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
  97. dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
  98. dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
  99. dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
  100. dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
  101. dw/server/ui/assets/css-DIMkf-bt.js +3 -0
  102. dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
  103. dw/server/ui/assets/cssMode-CPznxfY8.js +1 -0
  104. dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
  105. dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
  106. dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
  107. dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
  108. dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
  109. dw/server/ui/assets/editor.api-CpWcotrd.js +847 -0
  110. dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
  111. dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
  112. dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
  113. dw/server/ui/assets/freemarker2-CXtRM8N4.js +3 -0
  114. dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
  115. dw/server/ui/assets/go-C-y9NEjX.js +1 -0
  116. dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
  117. dw/server/ui/assets/handlebars-N7x-6NMY.js +1 -0
  118. dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
  119. dw/server/ui/assets/html-PhsdjHSr.js +1 -0
  120. dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
  121. dw/server/ui/assets/htmlMode-Dgj0SEok.js +1 -0
  122. dw/server/ui/assets/index-3Vw6WAPW.css +1 -0
  123. dw/server/ui/assets/index-DgrYhQd9.js +43 -0
  124. dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
  125. dw/server/ui/assets/java-BEtHBSE6.js +1 -0
  126. dw/server/ui/assets/javascript-BJqN9Qhv.js +1 -0
  127. dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
  128. dw/server/ui/assets/jsonMode-DbM4SWSv.js +7 -0
  129. dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
  130. dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
  131. dw/server/ui/assets/less-B9JPFI3C.js +2 -0
  132. dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
  133. dw/server/ui/assets/liquid-BWr8lEc4.js +1 -0
  134. dw/server/ui/assets/lspLanguageFeatures-C1iGuDyZ.js +4 -0
  135. dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
  136. dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
  137. dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
  138. dw/server/ui/assets/mdx-DAdMi_0p.js +1 -0
  139. dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
  140. dw/server/ui/assets/monaco--ixms01u.css +1 -0
  141. dw/server/ui/assets/monaco-BGCeEqaw.js +56 -0
  142. dw/server/ui/assets/msdax-DauUninz.js +1 -0
  143. dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
  144. dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
  145. dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
  146. dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
  147. dw/server/ui/assets/perl-oz_6vUea.js +1 -0
  148. dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
  149. dw/server/ui/assets/php-nr791fC2.js +1 -0
  150. dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
  151. dw/server/ui/assets/postiats-43DmfD33.js +1 -0
  152. dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
  153. dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
  154. dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
  155. dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
  156. dw/server/ui/assets/python-Bcn70HdC.js +1 -0
  157. dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
  158. dw/server/ui/assets/r-BwWrilGY.js +1 -0
  159. dw/server/ui/assets/razor-D1HmNnby.js +1 -0
  160. dw/server/ui/assets/redis-ClamHrr6.js +1 -0
  161. dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
  162. dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
  163. dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
  164. dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
  165. dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
  166. dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
  167. dw/server/ui/assets/scheme-BeGwcela.js +1 -0
  168. dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
  169. dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
  170. dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
  171. dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
  172. dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
  173. dw/server/ui/assets/sql-NEE52Syq.js +1 -0
  174. dw/server/ui/assets/st-DbInun42.js +1 -0
  175. dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
  176. dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
  177. dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
  178. dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
  179. dw/server/ui/assets/tsMode-D6u0XmOW.js +11 -0
  180. dw/server/ui/assets/twig-De2hgUGE.js +1 -0
  181. dw/server/ui/assets/typescript-BU6v-LMV.js +1 -0
  182. dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
  183. dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
  184. dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
  185. dw/server/ui/assets/workers-Cn7cTUKr.js +1 -0
  186. dw/server/ui/assets/xml--0LP2Lwk.js +1 -0
  187. dw/server/ui/assets/yaml-mpBg9jnt.js +1 -0
  188. dw/server/ui/index.html +17 -0
  189. dw/server/updater.py +192 -0
  190. dw/settings.py +98 -0
  191. dw/shot_span_preflight.py +116 -0
  192. dw/shots.py +359 -0
  193. dw/slice_preflight.py +148 -0
  194. dw/step.py +187 -0
  195. dw/step_cache.py +442 -0
  196. dw/subfolders.py +107 -0
  197. dw/task_domains.py +307 -0
  198. dw/tasks/assess.py +826 -0
  199. dw/tasks/audio_transcription.py +88 -0
  200. dw/tasks/audio_utils.py +1862 -0
  201. dw/tasks/background_remover.py +43 -0
  202. dw/tasks/borders.py +113 -0
  203. dw/tasks/compose_text.py +74 -0
  204. dw/tasks/concat_videos.py +300 -0
  205. dw/tasks/depth_estimator.py +54 -0
  206. dw/tasks/diffusion_upscale.py +109 -0
  207. dw/tasks/dissolve_videos.py +342 -0
  208. dw/tasks/format_messages.py +24 -0
  209. dw/tasks/gather.py +173 -0
  210. dw/tasks/grade.py +97 -0
  211. dw/tasks/image_to_text.py +43 -0
  212. dw/tasks/image_utils.py +764 -0
  213. dw/tasks/interpolate_frames.py +252 -0
  214. dw/tasks/judge.py +68 -0
  215. dw/tasks/model_cache.py +55 -0
  216. dw/tasks/pair_audio.py +268 -0
  217. dw/tasks/qr_code.py +19 -0
  218. dw/tasks/restore_faces.py +175 -0
  219. dw/tasks/rife_model.py +192 -0
  220. dw/tasks/segment.py +121 -0
  221. dw/tasks/select.py +111 -0
  222. dw/tasks/speech_generation.py +228 -0
  223. dw/tasks/stabilize.py +129 -0
  224. dw/tasks/task.py +920 -0
  225. dw/tasks/tensor_image.py +57 -0
  226. dw/tasks/text_generation.py +169 -0
  227. dw/tasks/text_sections.py +80 -0
  228. dw/tasks/upscale.py +203 -0
  229. dw/tasks/video_utils.py +624 -0
  230. dw/tasks/zoe_depth.py +71 -0
  231. dw/teacache.py +381 -0
  232. dw/teacache_models.json +99 -0
  233. dw/test.py +29 -0
  234. dw/type_helpers.py +231 -0
  235. dw/validate.py +68 -0
  236. dw/variable_constraints.py +444 -0
  237. dw/variables.py +443 -0
  238. dw/video_extensions.py +141 -0
  239. dw/vram_estimate.py +116 -0
  240. dw/worker.py +764 -0
  241. dw/workflow.py +2007 -0
  242. dw/workflow_schema.json +1346 -0
  243. dw/workflow_sources.py +383 -0
  244. dw/workflows/h3_context_ir.json +57 -0
  245. dw/workflows/test.json +31 -0
  246. dw/workspace.py +730 -0
  247. dw_mcp/__init__.py +6 -0
  248. dw_mcp/__main__.py +133 -0
  249. dw_mcp/assets.py +336 -0
  250. dw_mcp/authoring.py +114 -0
  251. dw_mcp/catalog.py +360 -0
  252. dw_mcp/client.py +486 -0
  253. dw_mcp/diagnose.py +371 -0
  254. dw_mcp/exports.py +84 -0
  255. dw_mcp/guides.py +35 -0
  256. dw_mcp/media.py +638 -0
  257. dw_mcp/models.py +97 -0
  258. dw_mcp/prompts.py +104 -0
  259. dw_mcp/server.py +1343 -0
  260. dw_mcp/workspaces.py +212 -0
@@ -0,0 +1,230 @@
1
+ """Whether a list-driven or composed run is projected to exceed host RAM.
2
+
3
+ `dw/host_memory.py` reads what a *running* worker is holding; this module is
4
+ the pre-flight half, for the failure mode `device_memory_stats`/VRAM warnings
5
+ never covered - a `for_each` that keeps its pipeline resident across a
6
+ 12-entry list, or a per-clip step whose footprint scales with the list, can
7
+ SIGKILL on host RAM with the accelerator nowhere near full (#243).
8
+
9
+ Scope is v1, by the repo owner's own sign-off on the issue: observed, not
10
+ curated (no author-declared memory figure - `host_memory_job_peak_rss_mb` is
11
+ a worker-reported field, not a schema key); warn, not refuse (host RAM headroom
12
+ is a property of *this machine*, not something a caller chose, so it never
13
+ blocks a run); no cross-machine normalization; and a cold start - no history
14
+ for this workflow at all - means no check, the same rule `observed_cost.py`
15
+ uses for GPU minutes.
16
+ """
17
+
18
+ import json
19
+ import logging
20
+
21
+ from .for_each import FOR_EACH_KEY
22
+ from .plan import _has_seedable_step
23
+
24
+ logger = logging.getLogger("dw")
25
+
26
+ # The fraction of physical RAM a projected peak may reach before this module
27
+ # says anything - matches nothing else in the codebase; there is no existing
28
+ # "safe headroom" constant to share, because nothing else projects against
29
+ # the machine's total RAM rather than a curated figure.
30
+ CEILING_FRACTION = 0.9
31
+
32
+ _RELEASE_FLAGS = ("release_pipeline", "release_models")
33
+
34
+
35
+ def releases_between_iterations(definition):
36
+ """Whether every `for_each` step in this workflow drops its pipeline (or
37
+ its models) between entries, rather than holding it resident for the
38
+ whole list.
39
+
40
+ `release_pipeline`/`release_models` survive on a for_each step's *last*
41
+ member only (`dw/for_each.py`), so their presence on the template step is
42
+ what says the author asked for the resident-vs-released shape here - the
43
+ same field, read for a different question. A workflow with no `for_each`
44
+ step at all answers True: there is no list to hold anything resident
45
+ across, so the "keeps everything resident" projection would not describe
46
+ it.
47
+ """
48
+ for_each_steps = [
49
+ step
50
+ for step in definition.get("steps") or []
51
+ if isinstance(step, dict) and FOR_EACH_KEY in step
52
+ ]
53
+ if not for_each_steps:
54
+ return True
55
+ return all(
56
+ any(step.get(flag) for flag in _RELEASE_FLAGS) for step in for_each_steps
57
+ )
58
+
59
+
60
+ def _requested_count(list_entries):
61
+ """The largest list length the caller's own arguments drive, or None
62
+ when this run has no list at all - such a run is never the failure mode
63
+ this module exists for."""
64
+ if not list_entries:
65
+ return None
66
+ counts = [value for value in list_entries.values() if isinstance(value, int)]
67
+ return max(counts) if counts else None
68
+
69
+
70
+ def _row_count(row, list_entries, variables):
71
+ """The list length a historical row ran with, read from its own stored
72
+ arguments rather than the current request's - a row that ran a shorter
73
+ or longer list is still comparable, once divided out, for the resident
74
+ projection's per-entry figure.
75
+
76
+ A row's stored `arguments` are only the caller's *overrides*
77
+ (`job.spec["arguments"]`), never merged with the workflow's declared
78
+ defaults - a run of the catalog default (the common case for this
79
+ module's own regression case) is recorded as `{}`. Falling back to
80
+ `variables[name]` the same way `observed_cost._bucket_key` does is what
81
+ lets such a row resolve to its actual list length rather than to no
82
+ length at all, which used to drop every default-arguments row out of the
83
+ per-entry figure and left the resident shape silent regardless of
84
+ history (#264).
85
+
86
+ Only the variable names the *current* request's `list_entries` names are
87
+ read back, on the assumption that a workflow's list-driving variables
88
+ are stable across its history - the same assumption `observed_cost.py`
89
+ makes bucketing by declared `cost_drivers`.
90
+ """
91
+ try:
92
+ arguments = json.loads(row.get("arguments") or "{}")
93
+ except (TypeError, ValueError):
94
+ arguments = {}
95
+ if not isinstance(arguments, dict):
96
+ arguments = {}
97
+ counts = []
98
+ for name in list_entries:
99
+ value = arguments.get(name, variables.get(name))
100
+ if isinstance(value, list):
101
+ counts.append(len(value))
102
+ return max(counts) if counts else None
103
+
104
+
105
+ def _peak_by_count(rows, list_entries, variables):
106
+ """This workflow's history grouped by list length, each count's peak
107
+ reduced to its median - the input a fixed+marginal fit reads."""
108
+ by_count = {}
109
+ for row in rows:
110
+ peak = row.get("host_memory_job_peak_rss_mb")
111
+ count = _row_count(row, list_entries, variables)
112
+ if isinstance(peak, (int, float)) and count:
113
+ by_count.setdefault(count, []).append(peak)
114
+ medians = {}
115
+ for count, peaks in by_count.items():
116
+ peaks.sort()
117
+ medians[count] = peaks[len(peaks) // 2]
118
+ return medians
119
+
120
+
121
+ def _fit_peak_model(medians):
122
+ """A run's peak as fixed (one model load) plus marginal (per extra
123
+ entry), fit from this workflow's own history rather than assumed.
124
+
125
+ A single model load dominates a resident run's peak - dividing the
126
+ whole observed figure by the entry count and multiplying back out
127
+ (the old per-entry model) counted that load once per entry instead of
128
+ once per run, projecting a 12-entry batch at roughly five times its
129
+ measured peak (#348). With history at two or more distinct list
130
+ lengths, the base and the marginal are fit through the smallest and
131
+ largest observed counts - the same "two extremes" a least-squares fit
132
+ would reduce to with only two points, and cheaper than one with more.
133
+ With history at only one list length, there is nothing to fit a slope
134
+ from: the measured peak is projected flat, which is the conservative
135
+ reading of "the measured behavior" the run has actually shown, rather
136
+ than guessing a growth rate with a single data point.
137
+
138
+ Returns `(base_mb, slope_mb_per_entry, low_count, high_count)`.
139
+ """
140
+ counts = sorted(medians)
141
+ if len(counts) == 1:
142
+ only = counts[0]
143
+ return medians[only], 0.0, only, only
144
+ low, high = counts[0], counts[-1]
145
+ slope = (medians[high] - medians[low]) / (high - low)
146
+ slope = max(slope, 0.0)
147
+ base = medians[low] - slope * low
148
+ return base, slope, low, high
149
+
150
+
151
+ def _already_survived(rows, list_entries, requested, projected_mb, variables):
152
+ """Whether a run at least as large as this request already finished on
153
+ this box at or above the projected peak.
154
+
155
+ A projection built from this workflow's own history can land above the
156
+ ceiling for a shape that is itself part of that history - the run being
157
+ projected already happened and succeeded (#254). `rows` only ever holds
158
+ successful runs (`JobHistory.finished_runs()`), so a matching row is a
159
+ completed data point, not a guess: warning about a peak already reached
160
+ without incident tells a caller nothing they don't already know from ten
161
+ minutes ago.
162
+ """
163
+ for row in rows:
164
+ peak = row.get("host_memory_job_peak_rss_mb")
165
+ if not isinstance(peak, (int, float)) or peak < projected_mb:
166
+ continue
167
+ count = _row_count(row, list_entries, variables)
168
+ if count is not None and count >= requested:
169
+ return True
170
+ return False
171
+
172
+
173
+ def host_memory_warnings(definition, list_entries, rows, ceiling_mb):
174
+ """Warnings for a projected host-memory peak this machine cannot hold,
175
+ or [] when there is nothing to project from or nothing to warn about.
176
+
177
+ `rows` are this workflow's finished runs as `JobHistory.finished_runs()`
178
+ groups them - the same history `observed_cost.py` reads, extended with
179
+ `host_memory_job_peak_rss_mb` per row (#243, repointed from the
180
+ process-lifetime `host_memory_peak_rss_mb` by #272 - a small job run
181
+ right after a heavy one no longer inherits the heavy job's peak).
182
+ `ceiling_mb` is this box's own RAM, scaled by `CEILING_FRACTION`; the
183
+ caller reads that once per request rather than this module importing
184
+ `host_memory` for a per-validate syscall.
185
+ """
186
+ requested = _requested_count(list_entries)
187
+ if requested is None or not rows or not ceiling_mb:
188
+ return []
189
+ if not _has_seedable_step(definition):
190
+ # A task-only workflow loads no model to accumulate across entries -
191
+ # the failure mode this module projects for doesn't apply to it
192
+ return []
193
+ variables = definition.get("variables") or {}
194
+ peaks = [row["host_memory_job_peak_rss_mb"] for row in rows]
195
+ peaks = [value for value in peaks if isinstance(value, (int, float))]
196
+ if not peaks:
197
+ # Cold start: history exists for this workflow, but no run of it
198
+ # ever reported a host-memory reading - nothing to project from
199
+ return []
200
+ if releases_between_iterations(definition):
201
+ projected_mb = max(peaks)
202
+ shape = f"~{round(projected_mb)} MB, the largest single iteration observed"
203
+ else:
204
+ medians = _peak_by_count(rows, list_entries, variables)
205
+ if not medians:
206
+ return []
207
+ base, slope, low_count, high_count = _fit_peak_model(medians)
208
+ projected_mb = base + slope * requested
209
+ if low_count == high_count:
210
+ shape = (
211
+ f"~{round(projected_mb)} MB at {requested} entries, based on "
212
+ f"this workflow's own {low_count}-entry runs with no larger "
213
+ "history yet to project growth from"
214
+ )
215
+ else:
216
+ shape = (
217
+ f"~{round(projected_mb)} MB at {requested} entries, "
218
+ f"extrapolated from runs of {low_count} and {high_count} entries"
219
+ )
220
+ if projected_mb <= ceiling_mb:
221
+ return []
222
+ if _already_survived(rows, list_entries, requested, projected_mb, variables):
223
+ return []
224
+ return [
225
+ f"Projected host memory for this run ({shape}) exceeds this "
226
+ f"machine's usable RAM (~{round(ceiling_mb)} MB) - based on "
227
+ f"{len(rows)} run(s) of this workflow's own history on this "
228
+ "machine, not a curated figure. The run is not blocked, but it "
229
+ "may be killed by the OS partway through."
230
+ ]
dw/hub_cache.py ADDED
@@ -0,0 +1,432 @@
1
+ """Inventory and deletion for the Hugging Face hub cache.
2
+
3
+ Every from_pretrained download lands in the hub cache and stays there until
4
+ something deletes it - a few video models is a few hundred GB. This wraps
5
+ huggingface_hub's own cache scanner so the server can show what is on disk
6
+ and free it, without inventing any path handling of its own: deletion goes
7
+ through scan_cache_dir's delete_revisions strategy, which only ever removes
8
+ revisions it found inside the cache directory.
9
+ """
10
+
11
+ import json
12
+ import logging
13
+ import os
14
+ import shutil
15
+ import threading
16
+ import time
17
+ import uuid
18
+
19
+ from huggingface_hub import constants, scan_cache_dir
20
+ from huggingface_hub.utils import CacheNotFound
21
+
22
+ from .security import SecurityError, validate_path
23
+
24
+ try:
25
+ # Xet-backed downloads aggregate into two bars built from our tracker
26
+ # class: reconstruction (file bytes written) and transfer (network
27
+ # bytes). Only reconstruction matches the file-size total the manager
28
+ # reports against - counting both would double the progress - and the
29
+ # transfer bar is recognizable by the bar_format it is created with
30
+ from huggingface_hub.utils._xet_progress_reporting import XET_TRANSFER_BAR_FORMAT
31
+ except ImportError: # older huggingface_hub - no xet aggregation
32
+ XET_TRANSFER_BAR_FORMAT = None
33
+
34
+ logger = logging.getLogger("dw")
35
+
36
+
37
+ def _resolved_cache_dir(cache_dir):
38
+ return str(cache_dir) if cache_dir else constants.HF_HUB_CACHE
39
+
40
+
41
+ def scan_models(cache_dir=None):
42
+ """The cache's contents as plain data: repos sorted largest-first,
43
+ with per-revision detail, plus disk totals for the volume it lives on."""
44
+ resolved = _resolved_cache_dir(cache_dir)
45
+ try:
46
+ scan = scan_cache_dir(resolved)
47
+ except CacheNotFound:
48
+ # No cache directory yet - nothing downloaded is a state, not an error
49
+ return {
50
+ "cache_dir": resolved,
51
+ "size_on_disk": 0,
52
+ "repos": [],
53
+ "warnings": [],
54
+ "disk_free": None,
55
+ "disk_total": None,
56
+ }
57
+
58
+ repos = []
59
+ for repo in scan.repos:
60
+ revisions = sorted(
61
+ repo.revisions, key=lambda rev: rev.last_modified or 0, reverse=True
62
+ )
63
+ repos.append(
64
+ {
65
+ "repo_id": repo.repo_id,
66
+ "repo_type": repo.repo_type,
67
+ "size_on_disk": repo.size_on_disk,
68
+ "nb_files": repo.nb_files,
69
+ "last_accessed": repo.last_accessed,
70
+ "last_modified": repo.last_modified,
71
+ "revisions": [
72
+ {
73
+ "commit_hash": revision.commit_hash,
74
+ "size_on_disk": revision.size_on_disk,
75
+ "refs": sorted(revision.refs),
76
+ "last_modified": revision.last_modified,
77
+ }
78
+ for revision in revisions
79
+ ],
80
+ }
81
+ )
82
+ repos.sort(key=lambda entry: entry["size_on_disk"], reverse=True)
83
+
84
+ usage = shutil.disk_usage(resolved)
85
+ return {
86
+ "cache_dir": resolved,
87
+ "size_on_disk": scan.size_on_disk,
88
+ "repos": repos,
89
+ "warnings": [str(warning) for warning in scan.warnings],
90
+ "disk_free": usage.free,
91
+ "disk_total": usage.total,
92
+ }
93
+
94
+
95
+ _WEIGHT_SUFFIXES = (".safetensors", ".bin", ".msgpack", ".onnx", ".pt")
96
+
97
+
98
+ def _repo_folder_name(repo_id, repo_type="model"):
99
+ return f"{repo_type}s--{repo_id.replace('/', '--')}"
100
+
101
+
102
+ def _snapshot_dir(repo_dir):
103
+ """The snapshot directory `refs/main` currently points at, or None -
104
+ a repo dir with no readable ref or no matching snapshot folder is
105
+ itself a sign the pull never finished."""
106
+ try:
107
+ with open(os.path.join(repo_dir, "refs", "main")) as f:
108
+ commit = f.read().strip()
109
+ except OSError:
110
+ return None
111
+ snapshot = os.path.join(repo_dir, "snapshots", commit)
112
+ return snapshot if os.path.isdir(snapshot) else None
113
+
114
+
115
+ def _component_incomplete(folder, variant):
116
+ """Whether a diffusers pipeline component's folder is missing files a
117
+ load would need. A component with no weight files at all (a scheduler
118
+ or tokenizer config) is complete once its folder exists; one with
119
+ weight files is checked against the requested `variant` only when that
120
+ variant is actually in play, since an unvarianted component (most
121
+ schedulers, safety checkers) never carries a tagged file."""
122
+ if not os.path.isdir(folder):
123
+ return True
124
+ try:
125
+ entries = os.listdir(folder)
126
+ except OSError:
127
+ return True
128
+ if not entries:
129
+ return True
130
+ if not variant:
131
+ return False
132
+ weight_files = [e for e in entries if e.endswith(_WEIGHT_SUFFIXES)]
133
+ if not weight_files:
134
+ return False
135
+ return not any(f".{variant}." in name for name in weight_files)
136
+
137
+
138
+ def repo_download_incomplete(repo_id, cache_dir=None, variant=None):
139
+ """Whether repo_id, though listed by `scan_models`, is left over from an
140
+ interrupted pull rather than fully fetched (#382): a cancelled
141
+ `snapshot_download` leaves the revision folder in place - the repo still
142
+ shows up in `scan_cache_dir` - but with an in-progress blob or with whole
143
+ components never fetched.
144
+
145
+ Two checks, in order of how cheap they are to rule out: any
146
+ `*.incomplete` blob (huggingface_hub's own marker for a file mid-transfer)
147
+ anywhere in the repo's `blobs/` makes the repo incomplete outright. Past
148
+ that, a diffusers pipeline repo (one with a `model_index.json` in its
149
+ current snapshot) is checked component by component, for the requested
150
+ `variant`; a plain checkpoint/LoRA repo (no `model_index.json`) has
151
+ nothing further to check beyond its blobs. This is not a full hub
152
+ file-list diff - it does not fetch the repo's file list from the hub, so
153
+ a file that was never attempted at all on an otherwise-untouched repo
154
+ (nothing downloaded, no blobs, no snapshot) is caught by the "no snapshot"
155
+ case, not enumerated.
156
+ """
157
+ resolved = _resolved_cache_dir(cache_dir)
158
+ try:
159
+ # _is_repo_id has already held repo_id to Hub's one-segment shape;
160
+ # this keeps the cache folder inside the cache on its own terms
161
+ repo_dir = validate_path(
162
+ os.path.join(resolved, _repo_folder_name(repo_id)), resolved
163
+ )
164
+ except SecurityError:
165
+ return True
166
+ if not os.path.isdir(repo_dir):
167
+ return True
168
+
169
+ blobs_dir = os.path.join(repo_dir, "blobs")
170
+ if os.path.isdir(blobs_dir):
171
+ try:
172
+ if any(name.endswith(".incomplete") for name in os.listdir(blobs_dir)):
173
+ return True
174
+ except OSError:
175
+ return True
176
+
177
+ snapshot = _snapshot_dir(repo_dir)
178
+ if snapshot is None:
179
+ return True
180
+
181
+ index_path = os.path.join(snapshot, "model_index.json")
182
+ if not os.path.isfile(index_path):
183
+ # Not a diffusers pipeline repo - blob completeness is all there is
184
+ return False
185
+ try:
186
+ with open(index_path) as f:
187
+ index = json.load(f)
188
+ except (OSError, ValueError):
189
+ return False
190
+
191
+ for key, value in index.items():
192
+ if key.startswith("_") or not isinstance(value, list):
193
+ continue
194
+ # A component is a folder beside model_index.json. The file comes
195
+ # from whoever published the repo, so an absolute or '..' key would
196
+ # point this listdir anywhere and make validate_workflow's
197
+ # downloads_required a yes/no oracle for directories on the box
198
+ if os.path.basename(key) != key or key in (".", ".."):
199
+ continue
200
+ if _component_incomplete(os.path.join(snapshot, key), variant):
201
+ return True
202
+ return False
203
+
204
+
205
+ def delete_model(repo_id, cache_dir=None):
206
+ """Delete every cached revision of repo_id. Returns the bytes freed.
207
+
208
+ Raises ValueError when the repo is not in the cache - the caller typed
209
+ or raced something, and nothing was deleted.
210
+ """
211
+ scan = scan_cache_dir(_resolved_cache_dir(cache_dir))
212
+ repo = next((r for r in scan.repos if r.repo_id == repo_id), None)
213
+ if repo is None:
214
+ raise ValueError(f"'{repo_id}' is not in the hub cache")
215
+
216
+ strategy = scan.delete_revisions(
217
+ *[revision.commit_hash for revision in repo.revisions]
218
+ )
219
+ freed = strategy.expected_freed_size
220
+ strategy.execute()
221
+ return freed
222
+
223
+
224
+ # ------------------------------------------------------------------ downloads
225
+
226
+
227
+ class DownloadCancelled(Exception):
228
+ """Raised inside a download's progress callback to abort it."""
229
+
230
+
231
+ class DownloadManager:
232
+ """Background snapshot downloads into the hub cache, with progress.
233
+
234
+ One thread per download; progress is fed by a tqdm-compatible tracker
235
+ that snapshot_download instantiates per file, so the counters aggregate
236
+ across the file pool. Cancellation raises out of the next progress tick;
237
+ huggingface_hub's partial files remain resumable, so a cancelled or
238
+ failed download picks up where it stopped when retried.
239
+ """
240
+
241
+ KEEP_FINISHED = 20
242
+
243
+ def __init__(self, download_fn=None, info_fn=None):
244
+ # Injectable for tests - the defaults reach the network
245
+ if download_fn is None or info_fn is None:
246
+ from huggingface_hub import HfApi, snapshot_download
247
+
248
+ download_fn = download_fn or snapshot_download
249
+ info_fn = info_fn or (
250
+ lambda repo_id: HfApi().repo_info(repo_id, files_metadata=True)
251
+ )
252
+ self._download_fn = download_fn
253
+ self._info_fn = info_fn
254
+ self._lock = threading.Lock()
255
+ self._downloads = {}
256
+
257
+ def start(self, repo_id):
258
+ """Begin downloading repo_id; returns the download's status dict.
259
+ Raises ValueError for an invalid repo id or one already in flight."""
260
+ from huggingface_hub.utils import HFValidationError, validate_repo_id
261
+
262
+ try:
263
+ validate_repo_id(repo_id)
264
+ except HFValidationError as e:
265
+ raise ValueError(str(e))
266
+
267
+ cancel_event = threading.Event()
268
+ with self._lock:
269
+ for entry in self._downloads.values():
270
+ if entry["repo_id"] == repo_id and entry["status"] == "downloading":
271
+ raise ValueError(f"'{repo_id}' is already downloading")
272
+ entry = {
273
+ "id": uuid.uuid4().hex[:12],
274
+ "repo_id": repo_id,
275
+ "status": "downloading",
276
+ "downloaded": 0,
277
+ "total": None,
278
+ "error": None,
279
+ "started_at": time.time(),
280
+ "finished_at": None,
281
+ # Set before the entry is published: cancel() reaches for this
282
+ # the moment the id is visible, and assigning it after the
283
+ # lock left a window where cancelling raised KeyError
284
+ "_cancel": cancel_event,
285
+ }
286
+ self._downloads[entry["id"]] = entry
287
+ self._prune()
288
+
289
+ thread = threading.Thread(
290
+ target=self._run, args=(entry, cancel_event), daemon=True
291
+ )
292
+ thread.start()
293
+ return self.status(entry["id"])
294
+
295
+ def _run(self, entry, cancel_event):
296
+ try:
297
+ try:
298
+ info = self._info_fn(entry["repo_id"])
299
+ total = sum(sibling.size or 0 for sibling in (info.siblings or []))
300
+ with self._lock:
301
+ entry["total"] = total or None
302
+ except Exception as e:
303
+ # Size is cosmetic; the download itself decides success
304
+ logger.debug(f"No size metadata for {entry['repo_id']}: {e}")
305
+
306
+ self._download_fn(
307
+ entry["repo_id"], tqdm_class=_tracker_class(self, entry, cancel_event)
308
+ )
309
+ self._finish(entry, "completed")
310
+ except DownloadCancelled:
311
+ self._finish(entry, "cancelled")
312
+ except Exception as e:
313
+ self._finish(entry, "failed", str(e))
314
+
315
+ def _finish(self, entry, status, error=None):
316
+ with self._lock:
317
+ entry["status"] = status
318
+ entry["error"] = error
319
+ entry["finished_at"] = time.time()
320
+
321
+ def _add_progress(self, entry, n):
322
+ with self._lock:
323
+ entry["downloaded"] += n
324
+
325
+ def cancel(self, download_id):
326
+ """Request cancellation; returns the status dict or None if unknown.
327
+ Takes effect at the download's next progress tick."""
328
+ with self._lock:
329
+ entry = self._downloads.get(download_id)
330
+ if entry is None:
331
+ return None
332
+ if entry["status"] == "downloading":
333
+ entry["_cancel"].set()
334
+ return self.status(download_id)
335
+
336
+ def status(self, download_id):
337
+ with self._lock:
338
+ entry = self._downloads.get(download_id)
339
+ return _public(entry) if entry else None
340
+
341
+ def status_list(self):
342
+ """Every tracked download, newest first."""
343
+ with self._lock:
344
+ entries = sorted(
345
+ self._downloads.values(),
346
+ key=lambda e: e["started_at"],
347
+ reverse=True,
348
+ )
349
+ return [_public(entry) for entry in entries]
350
+
351
+ def is_active(self):
352
+ with self._lock:
353
+ return any(
354
+ entry["status"] == "downloading" for entry in self._downloads.values()
355
+ )
356
+
357
+ def _prune(self):
358
+ # Called with the lock held: drop the oldest finished entries
359
+ finished = sorted(
360
+ (e for e in self._downloads.values() if e["status"] != "downloading"),
361
+ key=lambda e: e["started_at"],
362
+ )
363
+ excess = len(finished) - self.KEEP_FINISHED
364
+ for entry in finished[: max(0, excess)]:
365
+ del self._downloads[entry["id"]]
366
+
367
+
368
+ def _public(entry):
369
+ return {key: value for key, value in entry.items() if not key.startswith("_")}
370
+
371
+
372
+ def _tracker_class(manager, entry, cancel_event):
373
+ """A tqdm stand-in snapshot_download instantiates per file; every update
374
+ feeds the shared counters and honours cancellation."""
375
+
376
+ class Tracker:
377
+ def __init__(self, *args, **kwargs):
378
+ self.n = 0
379
+ self.total = kwargs.get("total")
380
+ # The xet transfer bar still ticks (cancellation works through
381
+ # it) but its network bytes stay out of the shared counters
382
+ self._counted = (
383
+ XET_TRANSFER_BAR_FORMAT is None
384
+ or kwargs.get("bar_format") != XET_TRANSFER_BAR_FORMAT
385
+ )
386
+
387
+ def update(self, n=1):
388
+ if cancel_event.is_set():
389
+ raise DownloadCancelled()
390
+ if n:
391
+ self.n += n
392
+ if self._counted:
393
+ manager._add_progress(entry, n)
394
+ return True
395
+
396
+ def close(self):
397
+ pass
398
+
399
+ def refresh(self):
400
+ pass
401
+
402
+ def set_description(self, *args, **kwargs):
403
+ pass
404
+
405
+ def set_postfix(self, *args, **kwargs):
406
+ pass
407
+
408
+ def set_postfix_str(self, *args, **kwargs):
409
+ pass
410
+
411
+ @property
412
+ def format_dict(self):
413
+ # What tqdm exposes for rate rendering; huggingface_hub reads
414
+ # 'rate' from it when composing the xet speed postfix
415
+ return {"n": self.n, "total": self.total, "elapsed": 0, "rate": None}
416
+
417
+ def __enter__(self):
418
+ return self
419
+
420
+ def __exit__(self, *exc):
421
+ self.close()
422
+ return False
423
+
424
+ @staticmethod
425
+ def get_lock():
426
+ return threading.RLock()
427
+
428
+ @staticmethod
429
+ def set_lock(lock):
430
+ pass
431
+
432
+ return Tracker