返回 JoyAI-Echo
helpers.py
1 import gc
2 import logging
3 from dataclasses import replace
4
5 import torch
6
7 from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderFactory
8 from ltx_core.components.noisers import Noiser
9 from ltx_core.components.protocols import DiffusionStepProtocol, GuiderProtocol
10 from ltx_core.conditioning import (
11 ConditioningItem,
12 VideoConditionByKeyframeIndex,
13 VideoConditionByLatentIndex,
14 )
15 from ltx_core.guidance.perturbations import (
16 BatchedPerturbationConfig,
17 Perturbation,
18 PerturbationConfig,
19 PerturbationType,
20 )
21 from ltx_core.model.transformer import Modality, X0Model
22 from ltx_core.model.video_vae import VideoEncoder
23 from ltx_core.text_encoders.gemma import GemmaTextEncoder
24 from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessorOutput
25 from ltx_core.tools import AudioLatentTools, LatentTools, VideoLatentTools
26 from ltx_core.types import AudioLatentShape, LatentState, VideoLatentShape, VideoPixelShape
27 from ltx_pipelines.utils.args import ImageConditioningInput
28 from ltx_pipelines.utils.media_io import decode_image, load_image_conditioning, resize_aspect_ratio_preserving
29 from ltx_pipelines.utils.types import (
30 DenoisingFunc,
31 DenoisingLoopFunc,
32 PipelineComponents,
33 )
34
35
36 def get_device() -> torch.device:
37 if torch.cuda.is_available():
38 return torch.device("cuda")
39 return torch.device("cpu")
40
41
42 def cleanup_memory() -> None:
43 gc.collect()
44 torch.cuda.empty_cache()
45 torch.cuda.synchronize()
46
47
48 def encode_prompts(
49 prompts: list[str],
50 model_ledger: object,
51 *,
52 enhance_prompt_image: str | None = None,
53 enhance_prompt_seed: int = 42,
54 enhance_first_prompt: bool = False,
55 ) -> list[EmbeddingsProcessorOutput]:
56 """Encode prompts through Gemma → embeddings processor, freeing each after use.
57 Loads the text encoder from *model_ledger*, optionally enhances the first
58 prompt, encodes all *prompts*, frees the text encoder, then loads the
59 embeddings processor to produce the final outputs. Because the text encoder
60 is loaded and freed entirely within this function, there are no lingering
61 references that could prevent GPU memory reclamation.
62 Args:
63 prompts: Text prompts to encode.
64 model_ledger: ModelLedger instance (used to load text encoder and embeddings processor).
65 enhance_prompt_image: Optional image path for prompt enhancement.
66 enhance_prompt_seed: Seed for prompt enhancement (default 42).
67 enhance_first_prompt: If True, enhance ``prompts[0]`` before encoding.
68 Returns:
69 List of EmbeddingsProcessorOutput, one per prompt.
70 """
71 text_encoder = model_ledger.text_encoder()
72 if enhance_first_prompt:
73 prompts = list(prompts)
74 prompts[0] = generate_enhanced_prompt(text_encoder, prompts[0], enhance_prompt_image, seed=enhance_prompt_seed)
75 raw_outputs = [text_encoder.encode(p) for p in prompts]
76 torch.cuda.synchronize()
77 del text_encoder
78 cleanup_memory()
79
80 embeddings_processor = model_ledger.gemma_embeddings_processor()
81 results: list[EmbeddingsProcessorOutput] = [
82 embeddings_processor.process_hidden_states(hs, mask) for hs, mask in raw_outputs
83 ]
84 del embeddings_processor
85 cleanup_memory()
86 return results
87
88
89 def combined_image_conditionings(
90 images: list[ImageConditioningInput],
91 height: int,
92 width: int,
93 video_encoder: VideoEncoder,
94 dtype: torch.dtype,
95 device: torch.device,
96 ) -> list[ConditioningItem]:
97 """Create a list of conditionings by replacing the latent at the first frame with the encoded image if present
98 and using other encoded images as the keyframe conditionings."""
99 conditionings = []
100 for img in images:
101 image = load_image_conditioning(
102 image_path=img.path,
103 height=height,
104 width=width,
105 dtype=dtype,
106 device=device,
107 crf=img.crf,
108 )
109 encoded_image = video_encoder(image)
110 if img.frame_idx == 0:
111 conditioning = VideoConditionByLatentIndex(
112 latent=encoded_image,
113 strength=img.strength,
114 latent_idx=0,
115 )
116 else:
117 conditioning = VideoConditionByKeyframeIndex(
118 keyframes=encoded_image,
119 strength=img.strength,
120 frame_idx=img.frame_idx,
121 )
122 conditionings.append(conditioning)
123 return conditionings
124
125
126 def image_conditionings_by_replacing_latent(
127 images: list[ImageConditioningInput],
128 height: int,
129 width: int,
130 video_encoder: VideoEncoder,
131 dtype: torch.dtype,
132 device: torch.device,
133 ) -> list[ConditioningItem]:
134 conditionings = []
135 for img in images:
136 image = load_image_conditioning(
137 image_path=img.path,
138 height=height,
139 width=width,
140 dtype=dtype,
141 device=device,
142 crf=img.crf,
143 )
144 encoded_image = video_encoder(image)
145 conditionings.append(
146 VideoConditionByLatentIndex(
147 latent=encoded_image,
148 strength=img.strength,
149 latent_idx=img.frame_idx,
150 )
151 )
152
153 return conditionings
154
155
156 def image_conditionings_by_adding_guiding_latent(
157 images: list[ImageConditioningInput],
158 height: int,
159 width: int,
160 video_encoder: VideoEncoder,
161 dtype: torch.dtype,
162 device: torch.device,
163 ) -> list[ConditioningItem]:
164 conditionings = []
165 for img in images:
166 image = load_image_conditioning(
167 image_path=img.path,
168 height=height,
169 width=width,
170 dtype=dtype,
171 device=device,
172 crf=img.crf,
173 )
174 encoded_image = video_encoder(image)
175 conditionings.append(
176 VideoConditionByKeyframeIndex(keyframes=encoded_image, frame_idx=img.frame_idx, strength=img.strength)
177 )
178 return conditionings
179
180
181 def noise_video_state(
182 output_shape: VideoPixelShape,
183 noiser: Noiser,
184 conditionings: list[ConditioningItem],
185 components: PipelineComponents,
186 dtype: torch.dtype,
187 device: torch.device,
188 noise_scale: float = 1.0,
189 initial_latent: torch.Tensor | None = None,
190 ) -> tuple[LatentState, VideoLatentTools]:
191 """Initialize and noise a video latent state for the diffusion pipeline.
192 Creates a video latent state from the output shape, applies conditionings,
193 and adds noise using the provided noiser. Returns the noised state and
194 video latent tools for further processing. If initial_latent is provided, it will be used to create the initial
195 state, otherwise an empty initial state will be created.
196 """
197 video_latent_shape = VideoLatentShape.from_pixel_shape(
198 shape=output_shape,
199 latent_channels=components.video_latent_channels,
200 scale_factors=components.video_scale_factors,
201 )
202 video_tools = VideoLatentTools(components.video_patchifier, video_latent_shape, output_shape.fps)
203 video_state = create_noised_state(
204 tools=video_tools,
205 conditionings=conditionings,
206 noiser=noiser,
207 dtype=dtype,
208 device=device,
209 noise_scale=noise_scale,
210 initial_latent=initial_latent,
211 )
212
213 return video_state, video_tools
214
215
216 def noise_audio_state(
217 output_shape: VideoPixelShape,
218 noiser: Noiser,
219 conditionings: list[ConditioningItem],
220 components: PipelineComponents,
221 dtype: torch.dtype,
222 device: torch.device,
223 noise_scale: float = 1.0,
224 initial_latent: torch.Tensor | None = None,
225 ) -> tuple[LatentState, AudioLatentTools]:
226 """Initialize and noise an audio latent state for the diffusion pipeline.
227 Creates an audio latent state from the output shape, applies conditionings,
228 and adds noise using the provided noiser. Returns the noised state and
229 audio latent tools for further processing. If initial_latent is provided, it will be used to create the initial
230 state, otherwise an empty initial state will be created.
231 """
232 audio_latent_shape = AudioLatentShape.from_video_pixel_shape(output_shape)
233 audio_tools = AudioLatentTools(components.audio_patchifier, audio_latent_shape)
234 audio_state = create_noised_state(
235 tools=audio_tools,
236 conditionings=conditionings,
237 noiser=noiser,
238 dtype=dtype,
239 device=device,
240 noise_scale=noise_scale,
241 initial_latent=initial_latent,
242 )
243
244 return audio_state, audio_tools
245
246
247 def create_noised_state(
248 tools: LatentTools,
249 conditionings: list[ConditioningItem],
250 noiser: Noiser,
251 dtype: torch.dtype,
252 device: torch.device,
253 noise_scale: float = 1.0,
254 initial_latent: torch.Tensor | None = None,
255 ) -> LatentState:
256 """Create a noised latent state from empty state, conditionings, and noiser.
257 Creates an empty latent state, applies conditionings, and then adds noise
258 using the provided noiser. Returns the final noised state ready for diffusion.
259 """
260 state = tools.create_initial_state(device, dtype, initial_latent)
261 state = state_with_conditionings(state, conditionings, tools)
262 state = noiser(state, noise_scale)
263
264 return state
265
266
267 def state_with_conditionings(
268 latent_state: LatentState, conditioning_items: list[ConditioningItem], latent_tools: LatentTools
269 ) -> LatentState:
270 """Apply a list of conditionings to a latent state.
271 Iterates through the conditioning items and applies each one to the latent
272 state in sequence. Returns the modified state with all conditionings applied.
273 """
274 for conditioning in conditioning_items:
275 latent_state = conditioning.apply_to(latent_state=latent_state, latent_tools=latent_tools)
276
277 return latent_state
278
279
280 def post_process_latent(denoised: torch.Tensor, denoise_mask: torch.Tensor, clean: torch.Tensor) -> torch.Tensor:
281 """Blend denoised output with clean state based on mask."""
282 return (denoised * denoise_mask + clean.float() * (1 - denoise_mask)).to(denoised.dtype)
283
284
285 def modality_from_latent_state(
286 state: LatentState,
287 context: torch.Tensor,
288 sigma: torch.Tensor,
289 enabled: bool = True,
290 ) -> Modality:
291 """Create a Modality from a latent state.
292 Constructs a Modality object with the latent state's data, timesteps derived
293 from the denoise mask and sigma, positions, and the provided context.
294 """
295 return Modality(
296 enabled=enabled,
297 latent=state.latent,
298 sigma=sigma,
299 timesteps=timesteps_from_mask(state.denoise_mask, sigma),
300 positions=state.positions,
301 context=context,
302 context_mask=None,
303 attention_mask=state.attention_mask,
304 )
305
306
307 def timesteps_from_mask(denoise_mask: torch.Tensor, sigma: float | torch.Tensor) -> torch.Tensor:
308 """Compute timesteps from a denoise mask and sigma value.
309 Multiplies the denoise mask by sigma to produce timesteps for each position
310 in the latent state. Areas where the mask is 0 will have zero timesteps.
311 """
312 return denoise_mask * sigma
313
314
315 def simple_denoising_func(
316 video_context: torch.Tensor, audio_context: torch.Tensor, transformer: X0Model
317 ) -> DenoisingFunc:
318 def simple_denoising_step(
319 video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
320 ) -> tuple[torch.Tensor, torch.Tensor]:
321 sigma = sigmas[step_index]
322 pos_video = modality_from_latent_state(video_state, video_context, sigma)
323 pos_audio = modality_from_latent_state(audio_state, audio_context, sigma)
324
325 denoised_video, denoised_audio = transformer(video=pos_video, audio=pos_audio, perturbations=None)
326 return denoised_video, denoised_audio
327
328 return simple_denoising_step
329
330
331 def guider_denoising_func(
332 guider: GuiderProtocol,
333 v_context_p: torch.Tensor,
334 v_context_n: torch.Tensor,
335 a_context_p: torch.Tensor,
336 a_context_n: torch.Tensor,
337 transformer: X0Model,
338 ) -> DenoisingFunc:
339 def guider_denoising_step(
340 video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
341 ) -> tuple[torch.Tensor, torch.Tensor]:
342 sigma = sigmas[step_index]
343 pos_video = modality_from_latent_state(video_state, v_context_p, sigma)
344 pos_audio = modality_from_latent_state(audio_state, a_context_p, sigma)
345
346 denoised_video, denoised_audio = transformer(video=pos_video, audio=pos_audio, perturbations=None)
347 if guider.enabled():
348 neg_video = modality_from_latent_state(video_state, v_context_n, sigma)
349 neg_audio = modality_from_latent_state(audio_state, a_context_n, sigma)
350
351 neg_denoised_video, neg_denoised_audio = transformer(video=neg_video, audio=neg_audio, perturbations=None)
352
353 denoised_video = denoised_video + guider.delta(denoised_video, neg_denoised_video)
354 denoised_audio = denoised_audio + guider.delta(denoised_audio, neg_denoised_audio)
355
356 return denoised_video, denoised_audio
357
358 return guider_denoising_step
359
360
361 def multi_modal_guider_denoising_func(
362 video_guider: MultiModalGuider,
363 audio_guider: MultiModalGuider,
364 v_context: torch.Tensor,
365 a_context: torch.Tensor,
366 transformer: X0Model,
367 *,
368 action_cond: dict | None = None,
369 last_denoised_video: torch.Tensor | None = None,
370 last_denoised_audio: torch.Tensor | None = None,
371 ) -> DenoisingFunc:
372 def guider_denoising_step(
373 video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
374 ) -> tuple[torch.Tensor, torch.Tensor]:
375 nonlocal last_denoised_video, last_denoised_audio
376
377 if video_guider.should_skip_step(step_index) and audio_guider.should_skip_step(step_index):
378 return last_denoised_video, last_denoised_audio
379
380 sigma = sigmas[step_index]
381 pos_video_modality = modality_from_latent_state(
382 video_state, v_context, sigma, enabled=not video_guider.should_skip_step(step_index)
383 )
384 pos_audio_modality = modality_from_latent_state(
385 audio_state, a_context, sigma, enabled=not audio_guider.should_skip_step(step_index)
386 )
387
388 denoised_video, denoised_audio = transformer(
389 video=pos_video_modality, audio=pos_audio_modality, perturbations=None, action_cond=action_cond
390 )
391 neg_denoised_video, neg_denoised_audio = 0.0, 0.0
392 if video_guider.do_unconditional_generation() or audio_guider.do_unconditional_generation():
393 if video_guider.do_unconditional_generation() and video_guider.negative_context is None:
394 raise ValueError("Negative context is required for unconditioned denoising")
395 if audio_guider.do_unconditional_generation() and audio_guider.negative_context is None:
396 raise ValueError("Negative context is required for unconditioned denoising")
397 neg_video_modality = modality_from_latent_state(
398 video_state,
399 video_guider.negative_context
400 if video_guider.negative_context is not None
401 else pos_video_modality.context,
402 sigma,
403 )
404 neg_audio_modality = modality_from_latent_state(
405 audio_state,
406 audio_guider.negative_context
407 if audio_guider.negative_context is not None
408 else pos_audio_modality.context,
409 sigma,
410 )
411
412 neg_denoised_video, neg_denoised_audio = transformer(
413 video=neg_video_modality, audio=neg_audio_modality, perturbations=None, action_cond=action_cond
414 )
415
416 ptb_denoised_video, ptb_denoised_audio = 0.0, 0.0
417 if video_guider.do_perturbed_generation() or audio_guider.do_perturbed_generation():
418 perturbations = []
419 if video_guider.do_perturbed_generation():
420 perturbations.append(
421 Perturbation(type=PerturbationType.SKIP_VIDEO_SELF_ATTN, blocks=video_guider.params.stg_blocks)
422 )
423 if audio_guider.do_perturbed_generation():
424 perturbations.append(
425 Perturbation(type=PerturbationType.SKIP_AUDIO_SELF_ATTN, blocks=audio_guider.params.stg_blocks)
426 )
427 perturbation_config = PerturbationConfig(perturbations=perturbations)
428 ptb_denoised_video, ptb_denoised_audio = transformer(
429 video=pos_video_modality, audio=pos_audio_modality,
430 perturbations=BatchedPerturbationConfig(perturbations=[perturbation_config]),
431 action_cond=action_cond,
432 )
433
434 mod_denoised_video, mod_denoised_audio = 0.0, 0.0
435 if video_guider.do_isolated_modality_generation() or audio_guider.do_isolated_modality_generation():
436 perturbations = [
437 Perturbation(type=PerturbationType.SKIP_A2V_CROSS_ATTN, blocks=None),
438 Perturbation(type=PerturbationType.SKIP_V2A_CROSS_ATTN, blocks=None),
439 ]
440 perturbation_config = PerturbationConfig(perturbations=perturbations)
441 mod_denoised_video, mod_denoised_audio = transformer(
442 video=pos_video_modality, audio=pos_audio_modality,
443 perturbations=BatchedPerturbationConfig(perturbations=[perturbation_config]),
444 action_cond=action_cond,
445 )
446
447 if video_guider.should_skip_step(step_index):
448 denoised_video = last_denoised_video
449 else:
450 denoised_video = video_guider.calculate(
451 denoised_video, neg_denoised_video, ptb_denoised_video, mod_denoised_video
452 )
453
454 if audio_guider.should_skip_step(step_index):
455 denoised_audio = last_denoised_audio
456 else:
457 denoised_audio = audio_guider.calculate(
458 denoised_audio, neg_denoised_audio, ptb_denoised_audio, mod_denoised_audio
459 )
460
461 last_denoised_video = denoised_video
462 last_denoised_audio = denoised_audio
463
464 return denoised_video, denoised_audio
465
466 return guider_denoising_step
467
468
469 def multi_modal_guider_factory_denoising_func(
470 video_guider_factory: MultiModalGuiderFactory,
471 audio_guider_factory: MultiModalGuiderFactory | None,
472 v_context: torch.Tensor,
473 a_context: torch.Tensor,
474 transformer: X0Model,
475 *,
476 action_cond: dict | None = None,
477 ) -> DenoisingFunc:
478 """Resolve guiders per step via factory.build_from_sigma, then multi_modal_guider_denoising_func."""
479 last_denoised_video: torch.Tensor | None = None
480 last_denoised_audio: torch.Tensor | None = None
481 sigma_vals_cached: list[float] | None = None
482
483 def guider_denoising_step(
484 video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
485 ) -> tuple[torch.Tensor, torch.Tensor]:
486 nonlocal last_denoised_video, last_denoised_audio, sigma_vals_cached
487 if sigma_vals_cached is None:
488 sigma_vals_cached = sigmas.detach().cpu().tolist()
489 sigma_val = sigma_vals_cached[step_index]
490 video_guider = video_guider_factory.build_from_sigma(sigma_val)
491 audio_guider = (audio_guider_factory or video_guider_factory).build_from_sigma(sigma_val)
492 denoise_fn = multi_modal_guider_denoising_func(
493 video_guider,
494 audio_guider,
495 v_context,
496 a_context,
497 transformer,
498 action_cond=action_cond,
499 last_denoised_video=last_denoised_video,
500 last_denoised_audio=last_denoised_audio,
501 )
502 denoised_video, denoised_audio = denoise_fn(video_state, audio_state, sigmas, step_index)
503 last_denoised_video, last_denoised_audio = denoised_video, denoised_audio
504 return denoised_video, denoised_audio
505
506 return guider_denoising_step
507
508
509 def denoise_audio_video( # noqa: PLR0913
510 output_shape: VideoPixelShape,
511 conditionings: list[ConditioningItem],
512 noiser: Noiser,
513 sigmas: torch.Tensor,
514 stepper: DiffusionStepProtocol,
515 denoising_loop_fn: DenoisingLoopFunc,
516 components: PipelineComponents,
517 dtype: torch.dtype,
518 device: torch.device,
519 noise_scale: float = 1.0,
520 initial_video_latent: torch.Tensor | None = None,
521 initial_audio_latent: torch.Tensor | None = None,
522 ) -> tuple[LatentState, LatentState]:
523 video_state, video_tools = noise_video_state(
524 output_shape=output_shape,
525 noiser=noiser,
526 conditionings=conditionings,
527 components=components,
528 dtype=dtype,
529 device=device,
530 noise_scale=noise_scale,
531 initial_latent=initial_video_latent,
532 )
533 audio_state, audio_tools = noise_audio_state(
534 output_shape=output_shape,
535 noiser=noiser,
536 conditionings=[],
537 components=components,
538 dtype=dtype,
539 device=device,
540 noise_scale=noise_scale,
541 initial_latent=initial_audio_latent,
542 )
543
544 video_state, audio_state = denoising_loop_fn(
545 sigmas,
546 video_state,
547 audio_state,
548 stepper,
549 )
550
551 video_state = video_tools.clear_conditioning(video_state)
552 video_state = video_tools.unpatchify(video_state)
553 audio_state = audio_tools.clear_conditioning(audio_state)
554 audio_state = audio_tools.unpatchify(audio_state)
555
556 return video_state, audio_state
557
558
559 def denoise_video_only( # noqa: PLR0913
560 output_shape: VideoPixelShape,
561 conditionings: list[ConditioningItem],
562 noiser: Noiser,
563 sigmas: torch.Tensor,
564 stepper: DiffusionStepProtocol,
565 denoising_loop_fn: DenoisingLoopFunc,
566 components: PipelineComponents,
567 dtype: torch.dtype,
568 device: torch.device,
569 noise_scale: float = 1.0,
570 initial_video_latent: torch.Tensor | None = None,
571 initial_audio_latent: torch.Tensor | None = None,
572 ) -> LatentState:
573 video_state, video_tools = noise_video_state(
574 output_shape=output_shape,
575 noiser=noiser,
576 conditionings=conditionings,
577 components=components,
578 dtype=dtype,
579 device=device,
580 noise_scale=noise_scale,
581 initial_latent=initial_video_latent,
582 )
583
584 audio_state, _ = noise_audio_state(
585 output_shape=output_shape,
586 noiser=noiser,
587 conditionings=[],
588 components=components,
589 dtype=dtype,
590 device=device,
591 noise_scale=0.0,
592 initial_latent=initial_audio_latent,
593 )
594
595 audio_state = replace(audio_state, denoise_mask=torch.zeros_like(audio_state.denoise_mask))
596
597 video_state, audio_state = denoising_loop_fn(
598 sigmas,
599 video_state,
600 audio_state,
601 stepper,
602 )
603
604 video_state = video_tools.clear_conditioning(video_state)
605 video_state = video_tools.unpatchify(video_state)
606
607 return video_state
608
609
610 _UNICODE_REPLACEMENTS = str.maketrans("\u2018\u2019\u201c\u201d\u2014\u2013\u00a0\u2032\u2212", "''\"\"-- '-")
611
612
613 def clean_response(text: str) -> str:
614 """Clean a response from curly quotes and leading non-letter characters which Gemma tends to insert."""
615 text = text.translate(_UNICODE_REPLACEMENTS)
616
617 # Remove leading non-letter characters
618 for i, char in enumerate(text):
619 if char.isalpha():
620 return text[i:]
621 return text
622
623
624 def generate_enhanced_prompt(
625 text_encoder: GemmaTextEncoder,
626 prompt: str,
627 image_path: str | None = None,
628 image_long_side: int = 896,
629 seed: int = 42,
630 ) -> str:
631 """Generate an enhanced prompt from a text encoder and a prompt."""
632 image = None
633 if image_path:
634 image = decode_image(image_path=image_path)
635 image = torch.tensor(image)
636 image = resize_aspect_ratio_preserving(image, image_long_side).to(torch.uint8)
637 prompt = text_encoder.enhance_i2v(prompt, image, seed=seed)
638 else:
639 prompt = text_encoder.enhance_t2v(prompt, seed=seed)
640 logging.info(f"Enhanced prompt: {prompt}")
641 return clean_response(prompt)
642
643
644 def assert_resolution(height: int, width: int, is_two_stage: bool) -> None:
645 """Assert that the resolution is divisible by the required divisor.
646 For two-stage pipelines, the resolution must be divisible by 64.
647 For one-stage pipelines, the resolution must be divisible by 32.
648 """
649 divisor = 64 if is_two_stage else 32
650 if height % divisor != 0 or width % divisor != 0:
651 raise ValueError(
652 f"Resolution ({height}x{width}) is not divisible by {divisor}. "
653 f"For {'two-stage' if is_two_stage else 'one-stage'} pipelines, "
654 f"height and width must be multiples of {divisor}."
655 )
656
656 lines PYTHON