diffusers-workflow 0.4.0a3__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.0a3.dist-info/METADATA +310 -0
- diffusers_workflow-0.4.0a3.dist-info/RECORD +171 -0
- diffusers_workflow-0.4.0a3.dist-info/WHEEL +5 -0
- diffusers_workflow-0.4.0a3.dist-info/entry_points.txt +6 -0
- diffusers_workflow-0.4.0a3.dist-info/licenses/LICENSE +201 -0
- diffusers_workflow-0.4.0a3.dist-info/top_level.txt +1 -0
- dw/__init__.py +353 -0
- dw/arguments.py +906 -0
- dw/cache_blocks.json +16 -0
- dw/cache_blocks.py +145 -0
- dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
- dw/events.py +78 -0
- dw/hub_cache.py +289 -0
- dw/introspection.py +458 -0
- dw/log_setup.py +45 -0
- dw/pipeline_processors/chain.py +750 -0
- dw/pipeline_processors/config_objects.py +235 -0
- dw/pipeline_processors/pipeline.py +1687 -0
- dw/pipeline_processors/remote.py +18 -0
- dw/previous_results.py +259 -0
- dw/prompt_weighting.py +378 -0
- dw/repl.py +298 -0
- dw/repl_commands.py +808 -0
- dw/repl_worker.py +129 -0
- dw/result.py +850 -0
- dw/run.py +92 -0
- dw/schema.py +24 -0
- dw/security.py +379 -0
- dw/serve.py +70 -0
- dw/server/__init__.py +2 -0
- dw/server/app.py +588 -0
- dw/server/jobs.py +547 -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-CEh6hWi2.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-CExg3_mM.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-DH6orYh2.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-CbrMVW4Q.js +1 -0
- dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
- dw/server/ui/assets/html-YDNPZw2M.js +1 -0
- dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
- dw/server/ui/assets/htmlMode-B_zSGWO2.js +1 -0
- dw/server/ui/assets/index-B7-VcYS-.css +1 -0
- dw/server/ui/assets/index-D_EiPU3b.js +13 -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-dYuBvioq.js +1 -0
- dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
- dw/server/ui/assets/jsonMode-CUqLM39V.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-D6vxBzMv.js +1 -0
- dw/server/ui/assets/lspLanguageFeatures-1WJ2palX.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-SHQb6vmD.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-CP-s5rcP.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-x0_EGHq9.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-BZC4LQDP.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-BTfA6SbD.js +11 -0
- dw/server/ui/assets/twig-De2hgUGE.js +1 -0
- dw/server/ui/assets/typescript-CWA4MsNk.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-CWU0uvj5.js +1 -0
- dw/server/ui/assets/xml-KmfTm3rg.js +1 -0
- dw/server/ui/assets/yaml-nFO_dDS6.js +1 -0
- dw/server/ui/index.html +17 -0
- dw/settings.py +77 -0
- dw/step.py +132 -0
- dw/tasks/audio_utils.py +266 -0
- dw/tasks/background_remover.py +43 -0
- dw/tasks/borders.py +113 -0
- dw/tasks/concat_videos.py +80 -0
- dw/tasks/depth_estimator.py +54 -0
- dw/tasks/diffusion_upscale.py +109 -0
- dw/tasks/format_messages.py +24 -0
- dw/tasks/gather.py +139 -0
- dw/tasks/image_to_text.py +43 -0
- dw/tasks/image_utils.py +661 -0
- dw/tasks/interpolate_frames.py +227 -0
- dw/tasks/model_cache.py +39 -0
- dw/tasks/pair_audio.py +58 -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/task.py +474 -0
- dw/tasks/tensor_image.py +57 -0
- dw/tasks/text_generation.py +168 -0
- dw/tasks/text_sections.py +80 -0
- dw/tasks/upscale.py +203 -0
- dw/tasks/video_utils.py +154 -0
- dw/tasks/zoe_depth.py +71 -0
- dw/teacache.py +376 -0
- dw/teacache_models.json +99 -0
- dw/test.py +29 -0
- dw/type_helpers.py +68 -0
- dw/validate.py +43 -0
- dw/variables.py +153 -0
- dw/worker.py +517 -0
- dw/workflow.py +553 -0
- dw/workflow_schema.json +1157 -0
- dw/workflows/augment_prompt.json +65 -0
- dw/workflows/describe_image.json +58 -0
- dw/workflows/h3_context_ir.json +57 -0
- dw/workflows/test.json +31 -0
dw/tasks/task.py
ADDED
|
@@ -0,0 +1,474 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
from typing import Callable, Dict
|
|
3
|
+
from .qr_code import get_qrcode_image
|
|
4
|
+
from .image_utils import process_image
|
|
5
|
+
from .video_utils import process_video
|
|
6
|
+
from .gather import gather_images, gather_inputs, gather_videos
|
|
7
|
+
from .format_messages import (
|
|
8
|
+
format_chat_message,
|
|
9
|
+
batch_decode_post_process,
|
|
10
|
+
get_dict_value,
|
|
11
|
+
)
|
|
12
|
+
|
|
13
|
+
# The model-backed handlers (upscale, restore_faces, segment, interpolate_frames,
|
|
14
|
+
# image_to_text, text_generation, diffusion_upscale) are imported inside their
|
|
15
|
+
# handlers - at module scope their transformers/model imports add seconds to
|
|
16
|
+
# every startup for workflows that never run those tasks
|
|
17
|
+
|
|
18
|
+
logger = logging.getLogger("dw")
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
# Command registry: maps command names to handler functions
|
|
22
|
+
_COMMAND_REGISTRY: Dict[str, Callable] = {}
|
|
23
|
+
|
|
24
|
+
# What each command's arguments actually are. The handlers forward
|
|
25
|
+
# **arguments into an implementation function, so that function's signature
|
|
26
|
+
# is the command's argument schema - registering its dotted path here (a
|
|
27
|
+
# string, to preserve the lazy-import discipline) lets the introspection
|
|
28
|
+
# layer read the same signature the runtime calls. 'provided' names the
|
|
29
|
+
# parameters the dispatch supplies itself, which are not workflow arguments.
|
|
30
|
+
_COMMAND_INFO: Dict[str, dict] = {}
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def register_command(command_name: str, implementation=None, provided=()):
|
|
34
|
+
"""
|
|
35
|
+
Decorator to register a command handler function.
|
|
36
|
+
|
|
37
|
+
Args:
|
|
38
|
+
command_name: The command name to register
|
|
39
|
+
implementation: Dotted path to the function whose signature defines
|
|
40
|
+
the command's arguments (None for a command that consumes a
|
|
41
|
+
free-form dict)
|
|
42
|
+
provided: Parameter names the dispatch supplies itself
|
|
43
|
+
|
|
44
|
+
Returns:
|
|
45
|
+
Decorator function
|
|
46
|
+
"""
|
|
47
|
+
|
|
48
|
+
def decorator(func: Callable) -> Callable:
|
|
49
|
+
_COMMAND_REGISTRY[command_name] = func
|
|
50
|
+
_COMMAND_INFO[command_name] = {
|
|
51
|
+
"kind": "command",
|
|
52
|
+
"implementation": implementation,
|
|
53
|
+
"provided": tuple(provided),
|
|
54
|
+
}
|
|
55
|
+
logger.debug(f"Registered command handler: {command_name}")
|
|
56
|
+
return func
|
|
57
|
+
|
|
58
|
+
return decorator
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def task_command_info(command_name):
|
|
62
|
+
"""Where a task command's argument schema lives: a dict with 'kind'
|
|
63
|
+
('command', 'image_processor' or 'video_processor'), 'implementation'
|
|
64
|
+
(dotted path or None for free-form), and 'provided'. Raises ValueError
|
|
65
|
+
for a name that is not a task command at all."""
|
|
66
|
+
info = _COMMAND_INFO.get(command_name)
|
|
67
|
+
if info is not None:
|
|
68
|
+
return info
|
|
69
|
+
if command_name in _VIDEO_PROCESSOR_INFO:
|
|
70
|
+
return _VIDEO_PROCESSOR_INFO[command_name]
|
|
71
|
+
from .image_utils import available_processors
|
|
72
|
+
|
|
73
|
+
if command_name in available_processors():
|
|
74
|
+
return {"kind": "image_processor", "implementation": None, "provided": ()}
|
|
75
|
+
raise ValueError(f"Unknown task command: '{command_name}'")
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
# Command handler functions
|
|
79
|
+
@register_command("qr_code", implementation="dw.tasks.qr_code.get_qrcode_image")
|
|
80
|
+
def _handle_qr_code(task, arguments, previous_pipelines):
|
|
81
|
+
"""Generate QR code image"""
|
|
82
|
+
logger.debug("Generating QR code")
|
|
83
|
+
return get_qrcode_image(**arguments)
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
@register_command("gather_images", implementation="dw.tasks.gather.gather_images")
|
|
87
|
+
def _handle_gather_images(task, arguments, previous_pipelines):
|
|
88
|
+
"""Gather multiple images"""
|
|
89
|
+
logger.debug("Gathering images")
|
|
90
|
+
return gather_images(**arguments)
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
@register_command("gather_videos", implementation="dw.tasks.gather.gather_videos")
|
|
94
|
+
def _handle_gather_videos(task, arguments, previous_pipelines):
|
|
95
|
+
"""Gather multiple videos"""
|
|
96
|
+
logger.debug("Gathering videos")
|
|
97
|
+
return gather_videos(**arguments)
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
# gather_inputs passes its whole dict through unchanged - free-form by design
|
|
101
|
+
@register_command("gather_inputs")
|
|
102
|
+
def _handle_gather_inputs(task, arguments, previous_pipelines):
|
|
103
|
+
"""Gather inputs from various sources"""
|
|
104
|
+
logger.debug("Gathering inputs")
|
|
105
|
+
return gather_inputs(arguments)
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
@register_command(
|
|
109
|
+
"concat_videos", implementation="dw.tasks.concat_videos.concat_videos"
|
|
110
|
+
)
|
|
111
|
+
def _handle_concat_videos(task, arguments, previous_pipelines):
|
|
112
|
+
"""Concatenate videos - and the audio generated with them - into one"""
|
|
113
|
+
logger.debug("Concatenating videos")
|
|
114
|
+
from .concat_videos import concat_videos
|
|
115
|
+
|
|
116
|
+
return concat_videos(**arguments)
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
@register_command("slice_audio", implementation="dw.tasks.audio_utils.slice_audio")
|
|
120
|
+
def _handle_slice_audio(task, arguments, previous_pipelines):
|
|
121
|
+
"""Cut a time- or frame-aligned slice out of an audio track"""
|
|
122
|
+
logger.debug("Slicing audio")
|
|
123
|
+
from .audio_utils import slice_audio
|
|
124
|
+
|
|
125
|
+
return slice_audio(**arguments)
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
@register_command("video_frames", implementation="dw.tasks.video_utils.frames_as_array")
|
|
129
|
+
def _handle_video_frames(task, arguments, previous_pipelines):
|
|
130
|
+
"""The frames of a generated video, as one array a later step can condition on"""
|
|
131
|
+
logger.debug("Extracting video frames")
|
|
132
|
+
from .video_utils import frames_as_array
|
|
133
|
+
|
|
134
|
+
return frames_as_array(**arguments)
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
@register_command("pair_audio", implementation="dw.tasks.pair_audio.pair_audio")
|
|
138
|
+
def _handle_pair_audio(task, arguments, previous_pipelines):
|
|
139
|
+
"""Pair a video's frames with an audio track generated beside them"""
|
|
140
|
+
logger.debug("Pairing audio with video")
|
|
141
|
+
from .pair_audio import pair_audio
|
|
142
|
+
|
|
143
|
+
return pair_audio(**arguments)
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
@register_command(
|
|
147
|
+
"crossfade_audio", implementation="dw.tasks.audio_utils.crossfade_audio"
|
|
148
|
+
)
|
|
149
|
+
def _handle_crossfade_audio(task, arguments, previous_pipelines):
|
|
150
|
+
"""Join audio tracks with an equal-power crossfade"""
|
|
151
|
+
logger.debug("Crossfading audio")
|
|
152
|
+
from .audio_utils import crossfade_audio
|
|
153
|
+
|
|
154
|
+
return crossfade_audio(**arguments)
|
|
155
|
+
|
|
156
|
+
|
|
157
|
+
@register_command(
|
|
158
|
+
"format_chat_message", implementation="dw.tasks.format_messages.format_chat_message"
|
|
159
|
+
)
|
|
160
|
+
def _handle_format_chat_message(task, arguments, previous_pipelines):
|
|
161
|
+
"""Format chat message for LLM input"""
|
|
162
|
+
logger.debug("Formatting chat message")
|
|
163
|
+
return format_chat_message(**arguments)
|
|
164
|
+
|
|
165
|
+
|
|
166
|
+
@register_command(
|
|
167
|
+
"get_dict_value", implementation="dw.tasks.format_messages.get_dict_value"
|
|
168
|
+
)
|
|
169
|
+
def _handle_get_dict_value(task, arguments, previous_pipelines):
|
|
170
|
+
"""Extract value from dictionary"""
|
|
171
|
+
logger.debug("Getting dictionary value")
|
|
172
|
+
return get_dict_value(**arguments)
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
@register_command("upscale", implementation="dw.tasks.upscale.upscale_image")
|
|
176
|
+
def _handle_upscale(task, arguments, previous_pipelines):
|
|
177
|
+
"""Upscale an image using a spandrel-compatible super-resolution model"""
|
|
178
|
+
logger.debug("Upscaling image")
|
|
179
|
+
image = arguments.pop("image")
|
|
180
|
+
model_name = arguments.pop("model_name")
|
|
181
|
+
from .upscale import upscale_image
|
|
182
|
+
|
|
183
|
+
return upscale_image(
|
|
184
|
+
image, model_name, device=task.device_for(arguments), **arguments
|
|
185
|
+
)
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
@register_command(
|
|
189
|
+
"diffusion_upscale", implementation="dw.tasks.diffusion_upscale.diffusion_upscale"
|
|
190
|
+
)
|
|
191
|
+
def _handle_diffusion_upscale(task, arguments, previous_pipelines):
|
|
192
|
+
"""Upscale an image using a diffusion-based upscale pipeline"""
|
|
193
|
+
logger.debug("Diffusion upscaling image")
|
|
194
|
+
image = arguments.pop("image")
|
|
195
|
+
from .diffusion_upscale import diffusion_upscale
|
|
196
|
+
|
|
197
|
+
return diffusion_upscale(image, device=task.device_for(arguments), **arguments)
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
@register_command(
|
|
201
|
+
"restore_faces", implementation="dw.tasks.restore_faces.restore_faces"
|
|
202
|
+
)
|
|
203
|
+
def _handle_restore_faces(task, arguments, previous_pipelines):
|
|
204
|
+
"""Restore faces in an image using a spandrel-compatible face restoration model"""
|
|
205
|
+
logger.debug("Restoring faces")
|
|
206
|
+
image = arguments.pop("image")
|
|
207
|
+
model_name = arguments.pop("model_name")
|
|
208
|
+
from .restore_faces import restore_faces
|
|
209
|
+
|
|
210
|
+
return restore_faces(
|
|
211
|
+
image, model_name, device=task.device_for(arguments), **arguments
|
|
212
|
+
)
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
@register_command("segment", implementation="dw.tasks.segment.segment_image")
|
|
216
|
+
def _handle_segment(task, arguments, previous_pipelines):
|
|
217
|
+
"""Segment objects in an image using text prompt"""
|
|
218
|
+
logger.debug("Segmenting image")
|
|
219
|
+
image = arguments.pop("image")
|
|
220
|
+
prompt = arguments.pop("prompt")
|
|
221
|
+
from .segment import segment_image
|
|
222
|
+
|
|
223
|
+
return segment_image(image, prompt, device=task.device_for(arguments), **arguments)
|
|
224
|
+
|
|
225
|
+
|
|
226
|
+
@register_command(
|
|
227
|
+
"interpolate_frames",
|
|
228
|
+
implementation="dw.tasks.interpolate_frames.interpolate_frames",
|
|
229
|
+
)
|
|
230
|
+
def _handle_interpolate_frames(task, arguments, previous_pipelines):
|
|
231
|
+
"""Interpolate video frames to increase frame rate"""
|
|
232
|
+
logger.debug("Interpolating frames")
|
|
233
|
+
video = arguments.pop("video")
|
|
234
|
+
from .interpolate_frames import interpolate_frames
|
|
235
|
+
|
|
236
|
+
return interpolate_frames(video, device=task.device_for(arguments), **arguments)
|
|
237
|
+
|
|
238
|
+
|
|
239
|
+
@register_command(
|
|
240
|
+
"image_to_text", implementation="dw.tasks.image_to_text.image_to_text"
|
|
241
|
+
)
|
|
242
|
+
def _handle_image_to_text(task, arguments, previous_pipelines):
|
|
243
|
+
"""Generate text caption from an image"""
|
|
244
|
+
logger.debug("Captioning image")
|
|
245
|
+
image = arguments.pop("image")
|
|
246
|
+
from .image_to_text import image_to_text
|
|
247
|
+
|
|
248
|
+
return image_to_text(image, device=task.device_for(arguments), **arguments)
|
|
249
|
+
|
|
250
|
+
|
|
251
|
+
@register_command(
|
|
252
|
+
"text_generation", implementation="dw.tasks.text_generation.generate_text"
|
|
253
|
+
)
|
|
254
|
+
def _handle_text_generation(task, arguments, previous_pipelines):
|
|
255
|
+
"""Generate text from a prompt using a local LLM"""
|
|
256
|
+
logger.debug("Generating text")
|
|
257
|
+
prompt = arguments.pop("prompt")
|
|
258
|
+
from .text_generation import generate_text
|
|
259
|
+
|
|
260
|
+
return generate_text(prompt, device=task.device_for(arguments), **arguments)
|
|
261
|
+
|
|
262
|
+
|
|
263
|
+
@register_command(
|
|
264
|
+
"extract_sections", implementation="dw.tasks.text_sections.extract_sections"
|
|
265
|
+
)
|
|
266
|
+
def _handle_extract_sections(task, arguments, previous_pipelines):
|
|
267
|
+
"""Reduce generated text to a known set of labelled sections"""
|
|
268
|
+
logger.debug("Extracting sections")
|
|
269
|
+
from .text_sections import extract_sections
|
|
270
|
+
|
|
271
|
+
return extract_sections(**arguments)
|
|
272
|
+
|
|
273
|
+
|
|
274
|
+
@register_command(
|
|
275
|
+
"batch_decode_post_process",
|
|
276
|
+
implementation="dw.tasks.format_messages.batch_decode_post_process",
|
|
277
|
+
provided=("processor",),
|
|
278
|
+
)
|
|
279
|
+
def _handle_batch_decode(task, arguments, previous_pipelines):
|
|
280
|
+
"""Batch decode post-processing with pipeline reference"""
|
|
281
|
+
logger.debug("Performing batch decode post-processing")
|
|
282
|
+
pipeline_reference = task.task_definition["pipeline_reference"]
|
|
283
|
+
if pipeline_reference not in previous_pipelines:
|
|
284
|
+
raise KeyError(
|
|
285
|
+
f"Pipeline reference '{pipeline_reference}' not found in previous pipelines. "
|
|
286
|
+
f"Available pipelines: {list(previous_pipelines.keys())}"
|
|
287
|
+
)
|
|
288
|
+
processor = previous_pipelines[pipeline_reference].pipeline
|
|
289
|
+
return batch_decode_post_process(processor, **arguments)
|
|
290
|
+
|
|
291
|
+
|
|
292
|
+
def _handle_image_processing(task, arguments, previous_pipelines):
|
|
293
|
+
"""Handle image processing commands"""
|
|
294
|
+
logger.debug("Processing image")
|
|
295
|
+
device = task.device_for(arguments)
|
|
296
|
+
return process_image(
|
|
297
|
+
arguments.pop("image"),
|
|
298
|
+
task.command,
|
|
299
|
+
device,
|
|
300
|
+
arguments,
|
|
301
|
+
)
|
|
302
|
+
|
|
303
|
+
|
|
304
|
+
def _handle_video_processing(task, arguments, previous_pipelines):
|
|
305
|
+
"""Handle video processing commands"""
|
|
306
|
+
logger.debug("Processing video")
|
|
307
|
+
device = task.device_for(arguments)
|
|
308
|
+
return process_video(
|
|
309
|
+
arguments.pop("video"),
|
|
310
|
+
task.command,
|
|
311
|
+
device,
|
|
312
|
+
arguments,
|
|
313
|
+
)
|
|
314
|
+
|
|
315
|
+
|
|
316
|
+
# Command names process_video (video_utils.py) accepts, with the function
|
|
317
|
+
# whose signature carries their arguments. video_utils dispatches via a plain
|
|
318
|
+
# if-chain, so keep this in sync with the branches in process_video().
|
|
319
|
+
# get_first/last_frame pin frame_index themselves, so it is 'provided'.
|
|
320
|
+
_VIDEO_PROCESSOR_INFO = {
|
|
321
|
+
"get_frame": {
|
|
322
|
+
"kind": "video_processor",
|
|
323
|
+
"implementation": "dw.tasks.video_utils.get_frame",
|
|
324
|
+
"provided": (),
|
|
325
|
+
},
|
|
326
|
+
"get_first_frame": {
|
|
327
|
+
"kind": "video_processor",
|
|
328
|
+
"implementation": "dw.tasks.video_utils.get_frame",
|
|
329
|
+
"provided": ("frame_index",),
|
|
330
|
+
},
|
|
331
|
+
"get_last_frame": {
|
|
332
|
+
"kind": "video_processor",
|
|
333
|
+
"implementation": "dw.tasks.video_utils.get_frame",
|
|
334
|
+
"provided": ("frame_index",),
|
|
335
|
+
},
|
|
336
|
+
}
|
|
337
|
+
_VIDEO_PROCESSOR_COMMANDS = sorted(_VIDEO_PROCESSOR_INFO)
|
|
338
|
+
|
|
339
|
+
|
|
340
|
+
class Task:
|
|
341
|
+
"""
|
|
342
|
+
Represents a task that can be executed as part of a workflow.
|
|
343
|
+
Tasks are atomic operations like image processing, data gathering, or message formatting.
|
|
344
|
+
"""
|
|
345
|
+
|
|
346
|
+
def __init__(self, task_definition, device):
|
|
347
|
+
"""
|
|
348
|
+
Initialize task with its configuration and device settings.
|
|
349
|
+
|
|
350
|
+
Args:
|
|
351
|
+
task_definition: Dictionary containing task configuration and parameters
|
|
352
|
+
device: Device to run task on (e.g., 'cuda', 'mps', 'cpu')
|
|
353
|
+
"""
|
|
354
|
+
self.task_definition = task_definition
|
|
355
|
+
self.device = device
|
|
356
|
+
logger.debug(f"Initialized task: {self.name} for device: {device}")
|
|
357
|
+
|
|
358
|
+
@property
|
|
359
|
+
def name(self):
|
|
360
|
+
"""Get task name from command property"""
|
|
361
|
+
return self.command
|
|
362
|
+
|
|
363
|
+
def device_for(self, arguments):
|
|
364
|
+
"""Get the device this task runs on, consuming any override in its arguments.
|
|
365
|
+
|
|
366
|
+
A task can pin itself to a device - a captioning model on the CPU while the GPU
|
|
367
|
+
holds a pipeline, for instance. The argument is removed either way so it does
|
|
368
|
+
not reach the command as a duplicate.
|
|
369
|
+
|
|
370
|
+
Args:
|
|
371
|
+
arguments: Arguments for this run of the task
|
|
372
|
+
|
|
373
|
+
Returns:
|
|
374
|
+
Device identifier the task should run on
|
|
375
|
+
"""
|
|
376
|
+
return arguments.pop("device", self.device)
|
|
377
|
+
|
|
378
|
+
@property
|
|
379
|
+
def argument_template(self):
|
|
380
|
+
"""
|
|
381
|
+
Get argument template for this task.
|
|
382
|
+
|
|
383
|
+
Returns:
|
|
384
|
+
Dictionary of arguments from inputs or arguments section
|
|
385
|
+
"""
|
|
386
|
+
# A task will either be an input array or a dictionary of arguments
|
|
387
|
+
if "inputs" in self.task_definition:
|
|
388
|
+
logger.debug("Using inputs as argument template")
|
|
389
|
+
return self.task_definition["inputs"]
|
|
390
|
+
|
|
391
|
+
logger.debug("Using arguments as argument template")
|
|
392
|
+
return self.task_definition["arguments"]
|
|
393
|
+
|
|
394
|
+
@property
|
|
395
|
+
def command(self):
|
|
396
|
+
"""Get command name or 'unknown' if not specified"""
|
|
397
|
+
return self.task_definition.get("command", "unknown")
|
|
398
|
+
|
|
399
|
+
def run(self, arguments, previous_pipelines={}):
|
|
400
|
+
"""
|
|
401
|
+
Execute the task with given arguments using the command registry.
|
|
402
|
+
|
|
403
|
+
Args:
|
|
404
|
+
arguments: Dictionary of arguments for task execution
|
|
405
|
+
previous_pipelines: Dictionary of previously created pipelines
|
|
406
|
+
|
|
407
|
+
Returns:
|
|
408
|
+
Task output based on command type
|
|
409
|
+
|
|
410
|
+
Raises:
|
|
411
|
+
ValueError: If command is unknown
|
|
412
|
+
KeyError: If required arguments or pipeline references are missing
|
|
413
|
+
"""
|
|
414
|
+
logger.debug(f"Running task: {self.command}")
|
|
415
|
+
logger.debug(f"Task arguments: {arguments}")
|
|
416
|
+
|
|
417
|
+
try:
|
|
418
|
+
# Cooperative cancellation reaches task steps too - without this
|
|
419
|
+
# a cancel during a long task waits for the whole task to finish
|
|
420
|
+
from ..events import get_context
|
|
421
|
+
|
|
422
|
+
get_context().check_cancelled()
|
|
423
|
+
|
|
424
|
+
# Look up command in registry
|
|
425
|
+
if self.command in _COMMAND_REGISTRY:
|
|
426
|
+
handler = _COMMAND_REGISTRY[self.command]
|
|
427
|
+
return handler(self, arguments, previous_pipelines)
|
|
428
|
+
|
|
429
|
+
# Not a registered command - check whether it names an image or
|
|
430
|
+
# video processor instead. Imported lazily here to preserve
|
|
431
|
+
# image_utils' lazy-import discipline for callers that never
|
|
432
|
+
# touch image processing.
|
|
433
|
+
from .image_utils import available_processors
|
|
434
|
+
|
|
435
|
+
if self.command in available_processors():
|
|
436
|
+
return _handle_image_processing(self, arguments, previous_pipelines)
|
|
437
|
+
|
|
438
|
+
if self.command in _VIDEO_PROCESSOR_COMMANDS:
|
|
439
|
+
return _handle_video_processing(self, arguments, previous_pipelines)
|
|
440
|
+
|
|
441
|
+
# Unknown command - not in the registry, and not a known image or
|
|
442
|
+
# video processor name either
|
|
443
|
+
error_msg = (
|
|
444
|
+
f"Unknown task command: '{self.command}'. "
|
|
445
|
+
f"Registered commands: {sorted(_COMMAND_REGISTRY.keys())}. "
|
|
446
|
+
f"Image processors: {available_processors()}. "
|
|
447
|
+
f"Video processors: {_VIDEO_PROCESSOR_COMMANDS}"
|
|
448
|
+
)
|
|
449
|
+
logger.error(error_msg)
|
|
450
|
+
raise ValueError(error_msg)
|
|
451
|
+
|
|
452
|
+
except KeyError as e:
|
|
453
|
+
# Missing required arguments or pipeline references
|
|
454
|
+
logger.error(
|
|
455
|
+
f"Missing required data for task {self.command}: {e}", exc_info=True
|
|
456
|
+
)
|
|
457
|
+
raise
|
|
458
|
+
except (ValueError, TypeError) as e:
|
|
459
|
+
# Invalid arguments or type mismatches
|
|
460
|
+
logger.error(
|
|
461
|
+
f"Invalid arguments for task {self.command}: {e}", exc_info=True
|
|
462
|
+
)
|
|
463
|
+
raise
|
|
464
|
+
except (OSError, IOError) as e:
|
|
465
|
+
# File operations, resource loading errors
|
|
466
|
+
logger.error(f"I/O error in task {self.command}: {e}", exc_info=True)
|
|
467
|
+
raise
|
|
468
|
+
except Exception as e:
|
|
469
|
+
# Catch-all for unexpected errors
|
|
470
|
+
logger.error(
|
|
471
|
+
f"Unexpected error ({type(e).__name__}) executing task {self.command}: {e}",
|
|
472
|
+
exc_info=True,
|
|
473
|
+
)
|
|
474
|
+
raise
|
dw/tasks/tensor_image.py
ADDED
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Shared PIL <-> float tensor conversions.
|
|
3
|
+
|
|
4
|
+
Consolidates the PIL-to-tensor round trip that was previously hand-rolled
|
|
5
|
+
independently in upscale.py, restore_faces.py, and interpolate_frames.py.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
import numpy as np
|
|
9
|
+
import torch
|
|
10
|
+
from PIL import Image
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def pil_to_float_tensor(image, device, dtype=None):
|
|
14
|
+
"""Convert a PIL image to a (1, 3, H, W) float tensor in [0, 1] on device.
|
|
15
|
+
|
|
16
|
+
The image is coerced to RGB first, so single-channel or RGBA inputs are
|
|
17
|
+
handled consistently. `dtype` defaults to float32 (the numpy source
|
|
18
|
+
precision); pass e.g. `torch.float16` to cast directly to a model's
|
|
19
|
+
working precision.
|
|
20
|
+
|
|
21
|
+
Args:
|
|
22
|
+
image: PIL Image
|
|
23
|
+
device: Target device (str or torch.device)
|
|
24
|
+
dtype: Optional torch dtype to cast to (default: float32)
|
|
25
|
+
|
|
26
|
+
Returns:
|
|
27
|
+
torch.Tensor of shape (1, 3, H, W)
|
|
28
|
+
"""
|
|
29
|
+
arr = np.array(image.convert("RGB")).astype(np.float32) / 255.0
|
|
30
|
+
tensor = torch.from_numpy(arr).permute(2, 0, 1).unsqueeze(0).to(device)
|
|
31
|
+
if dtype is not None:
|
|
32
|
+
tensor = tensor.to(dtype=dtype)
|
|
33
|
+
return tensor
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def float_tensor_to_pil(tensor):
|
|
37
|
+
"""Convert a (1, 3, H, W) or (3, H, W) float tensor in [0, 1] to a PIL RGB image.
|
|
38
|
+
|
|
39
|
+
Quantizes to uint8 via rounding (`.round()`) rather than truncation
|
|
40
|
+
(a bare truncating cast), matching diffusers' VaeImageProcessor.numpy_to_pil
|
|
41
|
+
behavior. This matters: with truncation, an exact 8-bit value that
|
|
42
|
+
round-trips through [0, 1] float (e.g. 128/255) can land a hair below
|
|
43
|
+
its integer (127.999...) and get chopped down a level instead of
|
|
44
|
+
landing back on 128. Rounding fixes that, at the cost of at most 1/255
|
|
45
|
+
of drift per channel versus the old truncating behavior for values that
|
|
46
|
+
were never exact to begin with.
|
|
47
|
+
|
|
48
|
+
Args:
|
|
49
|
+
tensor: torch.Tensor of shape (1, 3, H, W) or (3, H, W), values in [0, 1]
|
|
50
|
+
|
|
51
|
+
Returns:
|
|
52
|
+
PIL.Image.Image in RGB mode
|
|
53
|
+
"""
|
|
54
|
+
if tensor.dim() == 4:
|
|
55
|
+
tensor = tensor.squeeze(0)
|
|
56
|
+
arr = tensor.permute(1, 2, 0).mul(255).round().clamp(0, 255).byte().cpu().numpy()
|
|
57
|
+
return Image.fromarray(arr)
|
|
@@ -0,0 +1,168 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Text generation via HuggingFace transformers.
|
|
3
|
+
|
|
4
|
+
Takes a prompt (and optional system prompt) and generates text using a
|
|
5
|
+
local language model. Useful for prompt expansion, rewriting, and
|
|
6
|
+
other text-to-text tasks.
|
|
7
|
+
|
|
8
|
+
Supplying an image switches to a vision-language model, so the generated
|
|
9
|
+
text can describe what is actually in the picture rather than what the
|
|
10
|
+
prompt guesses is there. That is the difference between an image-conditioned
|
|
11
|
+
workflow whose prompt agrees with its keyframe and one whose prompt fights it.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
import logging
|
|
15
|
+
from transformers import pipeline as hf_pipeline
|
|
16
|
+
from .. import preferred_task_dtype
|
|
17
|
+
from .model_cache import cached_model
|
|
18
|
+
|
|
19
|
+
logger = logging.getLogger("dw")
|
|
20
|
+
|
|
21
|
+
_DEFAULT_MODEL = "Qwen/Qwen2.5-1.5B-Instruct"
|
|
22
|
+
# Small enough to stand in for the old captioning default; override for
|
|
23
|
+
# anything needing real detail
|
|
24
|
+
_DEFAULT_VISION_MODEL = "HuggingFaceTB/SmolVLM-256M-Instruct"
|
|
25
|
+
|
|
26
|
+
# Greedy decoding against a long, rigid format specification makes the vision
|
|
27
|
+
# models loop - finishing the answer, then repeating its closing sections until
|
|
28
|
+
# they run out of tokens. A penalty stops that without giving up reproducible
|
|
29
|
+
# output, which sampling would. Measured on Qwen3-VL against the H3 prompt
|
|
30
|
+
# spec: 1.05 still looped and filled the whole budget, 1.15 ended on its own at
|
|
31
|
+
# a length matching the format's own guidance. The text models do not need it
|
|
32
|
+
_VISION_REPETITION_PENALTY = 1.15
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def _image_part(image):
|
|
36
|
+
"""The chat content entry for an image.
|
|
37
|
+
|
|
38
|
+
A PIL image goes in as an object; a string is a URL or path the pipeline
|
|
39
|
+
loads itself, and it has to be declared as such - passing it under "image"
|
|
40
|
+
would hand the processor a bare string where it expects pixels.
|
|
41
|
+
"""
|
|
42
|
+
return {"type": "image", "url" if isinstance(image, str) else "image": image}
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def _build_messages(prompt, system_prompt, image):
|
|
46
|
+
"""Chat messages in the shape the chosen pipeline expects.
|
|
47
|
+
|
|
48
|
+
Text-generation models take plain string content. Vision models take a
|
|
49
|
+
list of typed parts, because the image is a part of the message rather
|
|
50
|
+
than something alongside it.
|
|
51
|
+
"""
|
|
52
|
+
messages = []
|
|
53
|
+
if image is None:
|
|
54
|
+
if system_prompt is not None:
|
|
55
|
+
messages.append({"role": "system", "content": system_prompt})
|
|
56
|
+
messages.append({"role": "user", "content": prompt})
|
|
57
|
+
else:
|
|
58
|
+
if system_prompt is not None:
|
|
59
|
+
messages.append(
|
|
60
|
+
{"role": "system", "content": [{"type": "text", "text": system_prompt}]}
|
|
61
|
+
)
|
|
62
|
+
messages.append(
|
|
63
|
+
{
|
|
64
|
+
"role": "user",
|
|
65
|
+
"content": [_image_part(image), {"type": "text", "text": prompt}],
|
|
66
|
+
}
|
|
67
|
+
)
|
|
68
|
+
return messages
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def generate_text(prompt, device="cpu", **kwargs):
|
|
72
|
+
"""Generate text from a prompt using a local language model.
|
|
73
|
+
|
|
74
|
+
Args:
|
|
75
|
+
prompt: The user message / prompt to expand or transform.
|
|
76
|
+
device: Target device ("cuda", "mps", "cpu").
|
|
77
|
+
**kwargs:
|
|
78
|
+
model_name: HuggingFace model ID. Defaults to
|
|
79
|
+
Qwen/Qwen2.5-1.5B-Instruct, or a small vision-language model
|
|
80
|
+
when an image is supplied. An image needs a model that can
|
|
81
|
+
accept one - a text-only model will fail to load as one.
|
|
82
|
+
system_prompt: Optional system instruction for the model.
|
|
83
|
+
max_new_tokens: Max tokens to generate (default: 500).
|
|
84
|
+
image: Optional PIL image, URL or path. Its presence is what
|
|
85
|
+
selects the vision pipeline.
|
|
86
|
+
repetition_penalty: Vision pipeline only (default: 1.15). Raise it
|
|
87
|
+
if a model still repeats itself, or set 1.0 to disable.
|
|
88
|
+
generate_kwargs: Anything else to hand the model's generate() -
|
|
89
|
+
no_repeat_ngram_size, top_p, min_new_tokens and so on. Merged
|
|
90
|
+
over what this function sets, so it can override those too.
|
|
91
|
+
|
|
92
|
+
Returns:
|
|
93
|
+
Generated text string.
|
|
94
|
+
"""
|
|
95
|
+
# A workflow declaring an optional image passes it through as null when the
|
|
96
|
+
# caller supplies none, and an empty string is the same statement
|
|
97
|
+
image = kwargs.get("image", None) or None
|
|
98
|
+
system_prompt = kwargs.get("system_prompt", None)
|
|
99
|
+
max_new_tokens = int(kwargs.get("max_new_tokens", 500))
|
|
100
|
+
|
|
101
|
+
if image is None:
|
|
102
|
+
pipeline_task = "text-generation"
|
|
103
|
+
model_name = kwargs.get("model_name", _DEFAULT_MODEL)
|
|
104
|
+
else:
|
|
105
|
+
pipeline_task = "image-text-to-text"
|
|
106
|
+
model_name = kwargs.get("model_name", _DEFAULT_VISION_MODEL)
|
|
107
|
+
|
|
108
|
+
dtype = preferred_task_dtype(device)
|
|
109
|
+
|
|
110
|
+
def load_pipe():
|
|
111
|
+
logger.info(f"Generating text with {model_name} on {device}")
|
|
112
|
+
return hf_pipeline(
|
|
113
|
+
pipeline_task,
|
|
114
|
+
model=model_name,
|
|
115
|
+
device_map=device,
|
|
116
|
+
torch_dtype=dtype,
|
|
117
|
+
)
|
|
118
|
+
|
|
119
|
+
# The task is part of the identity - the same model name can be loaded
|
|
120
|
+
# under either pipeline, and they are not interchangeable
|
|
121
|
+
pipe = cached_model(
|
|
122
|
+
("text_generation", pipeline_task, model_name, str(device), str(dtype)),
|
|
123
|
+
load_pipe,
|
|
124
|
+
)
|
|
125
|
+
|
|
126
|
+
messages = _build_messages(prompt, system_prompt, image)
|
|
127
|
+
|
|
128
|
+
# Decoding is greedy, so the same input returns the same text every run -
|
|
129
|
+
# which is what a workflow wants, and why there is nothing here to seed
|
|
130
|
+
generation = {"do_sample": False}
|
|
131
|
+
if image is not None:
|
|
132
|
+
generation["repetition_penalty"] = float(
|
|
133
|
+
kwargs.get("repetition_penalty", _VISION_REPETITION_PENALTY)
|
|
134
|
+
)
|
|
135
|
+
# Last, so a workflow can override anything decided above
|
|
136
|
+
generation.update(kwargs.get("generate_kwargs") or {})
|
|
137
|
+
|
|
138
|
+
return _generate(pipe, messages, image, max_new_tokens, generation)
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def _generate(pipe, messages, image, max_new_tokens, generation):
|
|
142
|
+
"""Run the pipeline, which each path calls differently."""
|
|
143
|
+
if image is None:
|
|
144
|
+
# This pipeline collects anything it does not name into the arguments it
|
|
145
|
+
# forwards to generate(), so settings go in as plain keywords
|
|
146
|
+
results = pipe(
|
|
147
|
+
messages,
|
|
148
|
+
max_new_tokens=max_new_tokens,
|
|
149
|
+
return_full_text=False,
|
|
150
|
+
**generation,
|
|
151
|
+
)
|
|
152
|
+
else:
|
|
153
|
+
# Images live inside the messages, so the chat goes in as `text` - the
|
|
154
|
+
# pipeline rejects a chat and an `images` argument together. Generation
|
|
155
|
+
# settings have to go through generate_kwargs here: anything else this
|
|
156
|
+
# pipeline does not name explicitly is forwarded to the processor and
|
|
157
|
+
# dropped, so a bare do_sample=False would leave sampling on. Passing
|
|
158
|
+
# max_new_tokens both ways is an error, so it stays a direct argument
|
|
159
|
+
results = pipe(
|
|
160
|
+
text=messages,
|
|
161
|
+
max_new_tokens=max_new_tokens,
|
|
162
|
+
return_full_text=False,
|
|
163
|
+
generate_kwargs=generation,
|
|
164
|
+
)
|
|
165
|
+
|
|
166
|
+
text = results[0]["generated_text"].strip()
|
|
167
|
+
logger.info(f"Generated: {text[:100]}{'...' if len(text) > 100 else ''}")
|
|
168
|
+
return text
|