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
|
@@ -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
|