diff --git a/CHANGELOG.md b/CHANGELOG.md
index 8fba412a9..4286e7dfe 100644
--- a/CHANGELOG.md
+++ b/CHANGELOG.md
@@ -1,11 +1,11 @@
# Change Log for SD.Next
-## Update for 2026-05-22
+## Update for 2026-05-23
-### Highlights for 2026-05-22
+### Highlights for 2026-05-23
*What's New?*
-- **Anima** made it to release version
+- **Anima** made it to release version, Microsoft joins the game with **Lens**
- **SDNQ** new quantization algorithm with even higher quality
- New **image analysis** feature and much improved **prompt enhance** capabilities which allow steering the model in real-time
- Improved image metadata options
@@ -14,11 +14,15 @@ And we have new [Contibuting** & **Development](https://vladmandic.github.io/sdn
Plus continued work on modernization of codebase: UI is now fully TypeScript based and new modular LoRA loader
-### Details for 2026-05-22
+### Details for 2026-05-23
- **Models**
- [CircleStone Anima 1.0](https://huggingface.co/circlestone-labs/Anima) in *Base* and *Turbo* (distilled) variants
in both original precision and SDNQ-4bit quantiztion
+ - [Microsoft Lens](https://huggingface.co/microsoft/Lens) in *Standard*, *Base* and *Turbo* (distilled) variant
+ 3.8B text-to-image DiT model with 12B GPT-OSS text-encoding and Flux2 VAE
+ oh, that 12B encoder is MoE with 3.6B activated plus its prequantized using `mxfp4`
+ *note* Lens comes with its own prompt-refiner, enable in settings -> model options (disabled by default)
- **Features**
- **SDNQ** new quantization algorithm: *Hadamard Rotations*
much higher quality than base SDNQ, but runs slightly slower
diff --git a/data/reference-distilled.json b/data/reference-distilled.json
index a371d8703..f1b6e7218 100644
--- a/data/reference-distilled.json
+++ b/data/reference-distilled.json
@@ -28,6 +28,15 @@
"tags": "distilled",
"extras": "sampler: Default, cfg_scale: 7.0"
},
+ "Microsoft Lens Turbo": {
+ "path": "microsoft/Lens-Turbo",
+ "preview": "microsoft--Lens-Turbo.jpg",
+ "desc": "Microsoft Lens-Turbo is the distilled Lens variant optimized for faster text-to-image generation with fewer steps.",
+ "skip": true,
+ "tags": "distilled",
+ "size": 30.53,
+ "date": "2026 May"
+ },
"Tencent FLUX.1 Dev SRPO": {
"path": "vladmandic/flux.1-dev-SRPO",
"preview": "vladmandic--flux.1-dev-SRPO.jpg",
diff --git a/data/reference.json b/data/reference.json
index 4e264beb9..89c95feff 100644
--- a/data/reference.json
+++ b/data/reference.json
@@ -38,6 +38,22 @@
"size": 6.94,
"date": "2023 July"
},
+ "Microsoft Lens": {
+ "path": "microsoft/Lens",
+ "preview": "microsoft--Lens.jpg",
+ "desc": "Microsoft Lens is a text-to-image DiT model using GPT-OSS chat-style prompt encoding and Flux2 VAE decoding.",
+ "skip": true,
+ "size": 30.53,
+ "date": "2026 May"
+ },
+ "Microsoft Lens Base": {
+ "path": "microsoft/Lens-Base",
+ "preview": "microsoft--Lens-Base.jpg",
+ "desc": "Microsoft Lens-Base is the base variant of Lens for text-to-image generation with GPT-OSS prompt features.",
+ "skip": true,
+ "size": 30.53,
+ "date": "2026 May"
+ },
"StabilityAI Stable Cascade": {
"path": "huggingface/stabilityai/stable-cascade",
"skip": true,
diff --git a/models/Reference/microsoft--Lens-Base.jpg b/models/Reference/microsoft--Lens-Base.jpg
new file mode 100644
index 000000000..e69de29bb
diff --git a/models/Reference/microsoft--Lens-Turbo.jpg b/models/Reference/microsoft--Lens-Turbo.jpg
new file mode 100644
index 000000000..e69de29bb
diff --git a/models/Reference/microsoft--Lens.jpg b/models/Reference/microsoft--Lens.jpg
new file mode 100644
index 000000000..e69de29bb
diff --git a/modules/modeldata.py b/modules/modeldata.py
index 4662f8165..4b1254828 100644
--- a/modules/modeldata.py
+++ b/modules/modeldata.py
@@ -94,6 +94,8 @@ def get_model_type(pipe):
model_type = 'kolors'
elif 'Meissonic' in name:
model_type = 'meissonic'
+ elif 'LensPipeline' in name:
+ model_type = 'lens'
elif 'Qwen' in name:
model_type = 'qwen'
elif 'ErnieImage' in name or 'ERNIE-Image' in name:
diff --git a/modules/sd_detect.py b/modules/sd_detect.py
index c575d80a0..14ba957c8 100644
--- a/modules/sd_detect.py
+++ b/modules/sd_detect.py
@@ -99,6 +99,8 @@ def guess_by_name(fn, current_guess):
new_guess = 'FLUX2 Klein'
elif 'flux.2' in fn.lower():
new_guess = 'FLUX2'
+ elif 'lens' in fn.lower():
+ new_guess = 'Lens'
elif 'ultraflux' in fn.lower():
new_guess = 'UltraFlux'
elif 'flux' in fn.lower() or 'flex.1' in fn.lower():
diff --git a/modules/sd_models.py b/modules/sd_models.py
index 51b61951d..5d7b4c741 100644
--- a/modules/sd_models.py
+++ b/modules/sd_models.py
@@ -486,6 +486,10 @@ def load_diffuser_force(detected_model_type, checkpoint_info, diffusers_load_con
from pipelines.model_hunyuandit import load_hunyuandit
sd_model = load_hunyuandit(checkpoint_info, diffusers_load_config)
allow_post_quant = False
+ elif model_type in ['Lens']:
+ from pipelines.model_lens import load_lens
+ sd_model = load_lens(checkpoint_info, diffusers_load_config)
+ allow_post_quant = False
elif model_type in ['Kandinsky 2.1']:
from pipelines.model_kandinsky import load_kandinsky21
sd_model = load_kandinsky21(checkpoint_info, diffusers_load_config)
diff --git a/modules/shared_items.py b/modules/shared_items.py
index 8cdd3d396..d642d3e16 100644
--- a/modules/shared_items.py
+++ b/modules/shared_items.py
@@ -54,6 +54,7 @@ pipelines = {
'ERNIE-Image': getattr(diffusers, 'ErnieImagePipeline', None),
'Nucleus-Image': getattr(diffusers, 'NucleusMoEImagePipeline', None),
'Z-Image': getattr(diffusers, 'ZImagePipeline', None),
+ 'Lens': getattr(diffusers, 'LensPipeline', None),
'FLUX2': getattr(diffusers, 'Flux2Pipeline', None),
'FLUX2 Klein': getattr(diffusers, 'Flux2KleinPipeline', None),
'LongCat': getattr(diffusers, 'LongCatImagePipeline', None),
@@ -142,6 +143,12 @@ def get_pipelines():
log.error(f'ONNX initialization error: {e}')
onnx_pipelines = {}
pipelines.update(onnx_pipelines)
+ if 'Lens' in pipelines and pipelines['Lens'] is None:
+ try:
+ import pipelines.lens as _lens
+ pipelines['Lens'] = getattr(diffusers, 'LensPipeline', None)
+ except Exception:
+ pass
for k, v in pipelines.items():
if k != 'Autodetect' and v is None:
from modules.logger import log
@@ -158,6 +165,8 @@ def get_repo(model):
return 'stabilityai/stable-diffusion-3.5-medium'
elif model == 'FluxPipeline' or model == 'FLUX':
return 'black-forest-labs/FLUX.1-dev'
+ elif model == 'LensPipeline' or model == 'Lens':
+ return 'microsoft/Lens'
else:
return None
diff --git a/modules/ui_definitions.py b/modules/ui_definitions.py
index 8fb3bc6f1..eb6cb14a7 100644
--- a/modules/ui_definitions.py
+++ b/modules/ui_definitions.py
@@ -113,6 +113,8 @@ def create_settings(cmd_opts):
"model_qwen_layers": OptionInfo(2, "Qwen layered number of layers", gr.Slider, {"minimum": 2, "maximum": 9, "step": 1 }),
"model_ernie_sep": OptionInfo("
ERNIE-Image
", "", gr.HTML),
"model_ernie_enable_pe": OptionInfo(False, "Enable prompt-enhance"),
+ "model_lens_sep": OptionInfo("Lens
", "", gr.HTML),
+ "model_lens_enable_pe": OptionInfo(False, "Enable prompt-enhance"),
}))
# --- Model Offloading ---
diff --git a/modules/vae/sd_vae_taesd.py b/modules/vae/sd_vae_taesd.py
index 6a733752d..46a18fd9d 100644
--- a/modules/vae/sd_vae_taesd.py
+++ b/modules/vae/sd_vae_taesd.py
@@ -71,7 +71,7 @@ def get_model(model_cls, variant=None):
elif model_cls in {'f1', 'h1', 'zimage', 'lumina2', 'chroma', 'longcat', 'omnigen2', 'flite', 'ovis', 'kandinsky5', 'glmimage', 'cogview3', 'cogview4', 'ultraflux'}:
model_cls = 'f1'
variant = 'TAE FLUX.1'
- elif model_cls in {'f2', 'ernieimage'}:
+ elif model_cls in {'f2', 'ernieimage', 'lens'}:
model_cls = 'f2'
variant = 'TAE FLUX.2'
elif model_cls in {'sd3'}:
diff --git a/pipelines/lens/__init__.py b/pipelines/lens/__init__.py
new file mode 100644
index 000000000..5254322c5
--- /dev/null
+++ b/pipelines/lens/__init__.py
@@ -0,0 +1,42 @@
+"""Lens - minimal text-to-image inference package."""
+
+import diffusers as _diffusers
+import transformers as _transformers
+from .pipeline import LensPipeline, LensPipelineOutput
+from .reasoner import PromptReasoner
+from .resolution import RESOLUTION_BUCKETS, resolve_resolution
+from .text_encoder import LensGptOssEncoder
+from .transformer import LensTransformer2DModel
+
+# ---------------------------------------------------------------------------
+# Make our custom subclasses discoverable by ``diffusers.DiffusionPipeline``.
+#
+# When ``LensPipeline.from_pretrained`` reads ``model_index.json``, it sees
+# entries like ``["transformers", "LensGptOssEncoder"]`` and runs
+# ``getattr(importlib.import_module("transformers"), "LensGptOssEncoder")``.
+# diffusers only allow-lists the libraries ``diffusers``, ``transformers`` and
+# ``onnxruntime.training`` - any other name is interpreted as a custom .py file
+# in the repo. So we inject our subclasses into those two namespaces here.
+#
+# Importing ``lens`` is required before calling ``LensPipeline.from_pretrained``
+# (this happens automatically when the user does ``from lens import LensPipeline``).
+# ---------------------------------------------------------------------------
+
+if not hasattr(_transformers, "LensGptOssEncoder"):
+ _transformers.LensGptOssEncoder = LensGptOssEncoder
+if not hasattr(_diffusers, "LensTransformer2DModel"):
+ _diffusers.LensTransformer2DModel = LensTransformer2DModel
+if not hasattr(_diffusers, "LensPipeline"):
+ _diffusers.LensPipeline = LensPipeline
+
+del _diffusers, _transformers
+
+__all__ = [
+ "LensPipeline",
+ "LensPipelineOutput",
+ "LensTransformer2DModel",
+ "LensGptOssEncoder",
+ "PromptReasoner",
+ "RESOLUTION_BUCKETS",
+ "resolve_resolution",
+]
diff --git a/pipelines/lens/pipeline.py b/pipelines/lens/pipeline.py
new file mode 100644
index 000000000..c167b5312
--- /dev/null
+++ b/pipelines/lens/pipeline.py
@@ -0,0 +1,647 @@
+"""Lens text-to-image pipeline.
+
+The pipeline follows the standard ``diffusers`` component and call conventions:
+components are registered via ``register_modules`` and the call signature
+supports ``height``/``width``, ``generator``, ``prompt_embeds``, ``output_type``,
+``return_dict``, and ``callback_on_step_end``.
+"""
+
+from __future__ import annotations
+
+from dataclasses import dataclass
+from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Sequence, Union
+
+import numpy as np
+import torch
+from diffusers import (
+ AutoencoderKLFlux2,
+ DiffusionPipeline,
+ FlowMatchEulerDiscreteScheduler,
+)
+from diffusers.utils import BaseOutput
+from diffusers.utils.torch_utils import randn_tensor
+from einops import rearrange
+from PIL import Image
+
+from .reasoner import PromptReasoner
+from .resolution import resolve_resolution
+
+if TYPE_CHECKING:
+ from transformers import PreTrainedTokenizerBase
+ from .text_encoder import LensGptOssEncoder
+ from .transformer import LensTransformer2DModel
+
+
+# ---------------------------------------------------------------------------
+# Helpers
+# ---------------------------------------------------------------------------
+
+
+def compute_empirical_mu(image_seq_len: int, num_steps: int) -> float:
+ """Empirical ``mu`` for ``FlowMatchEulerDiscreteScheduler`` dynamic shift.
+
+ Constants are calibrated for the Lens inference schedule.
+ """
+ a1, b1 = 8.73809524e-05, 1.89833333
+ a2, b2 = 0.00016927, 0.45666666
+ if image_seq_len > 4300:
+ return float(a2 * image_seq_len + b2)
+ m_200 = a2 * image_seq_len + b2
+ m_10 = a1 * image_seq_len + b1
+ a = (m_200 - m_10) / 190.0
+ b = m_200 - 200.0 * a
+ return float(a * num_steps + b)
+
+
+# Chat template constants used by the Lens text encoder.
+_CHAT_SYSTEM = (
+ "Describe the image by detailing the color, shape, size, texture, "
+ "quantity, text, spatial relationships of the objects and background."
+)
+_CHAT_ASSISTANT_THINKING = "Need to generate one image according to the description."
+DEFAULT_TXT_OFFSET = 97
+
+
+# Default Lens transformer architecture.
+DEFAULT_TRANSFORMER_CONFIG = dict(
+ patch_size=2,
+ in_channels=128,
+ out_channels=32,
+ num_layers=48,
+ attention_head_dim=64,
+ num_attention_heads=24,
+ inner_dim=1536,
+ enc_hidden_dim=2880,
+ axes_dims_rope=(8, 28, 28),
+ gate_mlp=True,
+ rms_norm=True,
+ multi_layer_encoder_feature=True,
+ selected_layer_index=(5, 11, 17, 23),
+)
+
+
+@dataclass
+class LensPipelineOutput(BaseOutput):
+ """Output of :class:`LensPipeline`.
+
+ Args:
+ images: list of decoded PIL images, or a numpy array of shape
+ ``[B, H, W, C]`` when ``output_type='np'``, or the raw latent
+ tensor when ``output_type='latent'``.
+ """
+
+ images: Union[List[Image.Image], np.ndarray, torch.Tensor]
+
+
+# ---------------------------------------------------------------------------
+# Pipeline
+# ---------------------------------------------------------------------------
+
+
+class LensPipeline(DiffusionPipeline):
+ r"""Lens text-to-image pipeline (GPT-OSS multi-layer features + Flux2 VAE).
+
+ Args:
+ scheduler ([`FlowMatchEulerDiscreteScheduler`]):
+ A scheduler used together with ``transformer`` to denoise the
+ encoded image latents.
+ vae ([`AutoencoderKLFlux2`]):
+ Flux2 VAE used to decode latents into pixel images.
+ text_encoder ([`LensGptOssEncoder`]):
+ ``GptOssForCausalLM`` subclass that exposes hidden states at the
+ configured ``selected_layer_index`` via ``encode_layers(...)``.
+ tokenizer ([`PreTrainedTokenizerBase`]):
+ GPT-OSS tokenizer.
+ transformer ([`LensTransformer2DModel`]):
+ The Lens denoising DiT.
+ reasoner ([`PromptReasoner`], *optional*):
+ Optional prompt rewriter (local OSS ``generate`` or
+ OpenAI-compatible API).
+ """
+
+ model_cpu_offload_seq = "text_encoder->transformer->vae"
+ _callback_tensor_inputs = [
+ "latents", "prompt_embeds", "negative_prompt_embeds",
+ ]
+
+ def __init__(
+ self,
+ scheduler: FlowMatchEulerDiscreteScheduler,
+ vae: AutoencoderKLFlux2,
+ text_encoder: LensGptOssEncoder,
+ tokenizer: PreTrainedTokenizerBase,
+ transformer: LensTransformer2DModel,
+ reasoner: Optional[PromptReasoner] = True,
+ ) -> None:
+ super().__init__()
+ self.register_modules(
+ scheduler=scheduler,
+ vae=vae,
+ text_encoder=text_encoder,
+ tokenizer=tokenizer,
+ transformer=transformer,
+ reasoner=reasoner,
+ )
+ if self.tokenizer.pad_token_id is None:
+ self.tokenizer.pad_token = self.tokenizer.eos_token
+ self.tokenizer.padding_side = "right"
+ # Flux2 latent tile factor (4x4 patchify) and Lens DiT in_channels=128.
+ self.vae_scale_factor = 16
+ self.latent_channels = self.transformer.config.in_channels
+ self.txt_offset = DEFAULT_TXT_OFFSET
+ self.default_sample_size = 1024
+
+ if not hasattr(self.text_encoder, "_lens_selected_layers"):
+ self.text_encoder.set_selected_layers(
+ self.transformer.config.selected_layer_index
+ )
+
+ if reasoner is not None:
+ self.reasoner = PromptReasoner(
+ text_encoder=self.text_encoder, tokenizer=self.tokenizer
+ )
+
+ # ------------------------------------------------------------------
+ # Prompt encoding
+ # ------------------------------------------------------------------
+
+ def _build_chat_inputs(
+ self, prompts: Sequence[str], max_sequence_length: int, device: torch.device
+ ):
+ rendered: List[str] = []
+ for prompt in prompts:
+ conversation = [
+ {"role": "system", "content": _CHAT_SYSTEM, "thinking": None},
+ {"role": "user", "content": prompt, "thinking": None},
+ {"role": "assistant", "thinking": _CHAT_ASSISTANT_THINKING, "content": ""},
+ ]
+ text = self.tokenizer.apply_chat_template(
+ conversation, tokenize=False, add_generation_prompt=False
+ )
+ text = text.split("<|return|>")[0]
+ rendered.append(text)
+
+ encoded = self.tokenizer(
+ rendered,
+ padding=True,
+ truncation=True,
+ max_length=max_sequence_length,
+ return_tensors="pt",
+ add_special_tokens=True,
+ )
+ return encoded["input_ids"].to(device), encoded["attention_mask"].to(device)
+
+ @torch.no_grad()
+ def _get_text_embeddings(
+ self, prompts: List[str], max_sequence_length: int, device: torch.device
+ ):
+ input_ids, attn_mask = self._build_chat_inputs(prompts, max_sequence_length, device)
+ layer_outputs = self.text_encoder.encode_layers(input_ids, attn_mask)
+
+ offset = self.txt_offset
+ if input_ids.shape[1] > offset:
+ features = [feat[:, offset:, :].contiguous() for feat in layer_outputs]
+ mask = attn_mask[:, offset:].bool()
+ else:
+ zero_shape = (input_ids.shape[0], 0, layer_outputs[0].shape[-1])
+ features = [layer_outputs[0].new_zeros(zero_shape) for _ in layer_outputs]
+ mask = torch.zeros(
+ (input_ids.shape[0], 0), dtype=torch.bool, device=device
+ )
+ return features, mask
+
+ def encode_prompt(
+ self,
+ prompt: Union[str, List[str]],
+ negative_prompt: Union[str, List[str]] = "",
+ num_images_per_prompt: int = 1,
+ prompt_embeds: Optional[List[torch.Tensor]] = None,
+ prompt_mask: Optional[torch.Tensor] = None,
+ negative_prompt_embeds: Optional[List[torch.Tensor]] = None,
+ negative_prompt_mask: Optional[torch.Tensor] = None,
+ max_sequence_length: int = 512,
+ device: Optional[torch.device] = None,
+ ):
+ """Encode positives and negatives. Returns
+ ``(prompt_embeds, prompt_mask, negative_prompt_embeds, negative_prompt_mask)``
+ where each ``*_embeds`` is a list of per-layer tensors and each
+ ``*_mask`` is a bool ``[B*N, S]`` tensor.
+
+ Each unique prompt is encoded **once**; the resulting features and mask
+ are then ``repeat_interleave``-d ``num_images_per_prompt`` times along
+ the batch axis. This preserves the ``[p0,p0,...,p1,p1,...]`` ordering
+ downstream consumers expect.
+ """
+ device = device or self._execution_device
+
+ prompts = [prompt] if isinstance(prompt, str) else list(prompt)
+ n = int(num_images_per_prompt)
+
+ # Negatives broadcast.
+ if isinstance(negative_prompt, str):
+ negatives = [negative_prompt] * len(prompts)
+ else:
+ negatives = list(negative_prompt)
+ if len(negatives) == 1:
+ negatives = negatives * len(prompts)
+ if len(negatives) != len(prompts):
+ raise ValueError(
+ "negative_prompt must be a string or a list of the same "
+ "length as prompt"
+ )
+
+ if prompt_embeds is None:
+ prompt_embeds, prompt_mask = self._get_text_embeddings(
+ prompts, max_sequence_length, device
+ )
+ prompt_embeds, prompt_mask = self._repeat_for_n(prompt_embeds, prompt_mask, n)
+ elif prompt_mask is None:
+ raise ValueError("`prompt_mask` must be provided when passing `prompt_embeds`.")
+ if negative_prompt_embeds is None:
+ if all(isinstance(neg, str) and not neg.strip() for neg in negatives):
+ # Empty negatives use an unconditional branch with no text tokens.
+ negative_prompt_embeds = [
+ feat.new_zeros(feat.shape) for feat in prompt_embeds
+ ]
+ negative_prompt_mask = torch.zeros_like(prompt_mask, dtype=torch.bool)
+ else:
+ negative_prompt_embeds, negative_prompt_mask = self._get_text_embeddings(
+ negatives, max_sequence_length, device
+ )
+ negative_prompt_embeds, negative_prompt_mask = self._repeat_for_n(
+ negative_prompt_embeds, negative_prompt_mask, n
+ )
+ elif negative_prompt_mask is None:
+ raise ValueError(
+ "`negative_prompt_mask` must be provided when passing "
+ "`negative_prompt_embeds`."
+ )
+ return prompt_embeds, prompt_mask, negative_prompt_embeds, negative_prompt_mask
+
+ @staticmethod
+ def _repeat_for_n(features: List[torch.Tensor], mask: torch.Tensor, n: int):
+ """Repeat each sample ``n`` times along the batch axis (interleaved)."""
+ if n == 1:
+ return features, mask
+ features = [f.repeat_interleave(n, dim=0) for f in features]
+ mask = mask.repeat_interleave(n, dim=0)
+ return features, mask
+
+ # ------------------------------------------------------------------
+ # Reasoner shim
+ # ------------------------------------------------------------------
+
+ def refine_prompt(
+ self, prompts: Sequence[str], enable_reasoner: bool = False
+ ) -> List[str]:
+ if self.reasoner is None:
+ return list(prompts)
+ print('HERE REFINE')
+ return self.reasoner.refine(prompts, enable=enable_reasoner)
+
+ # ------------------------------------------------------------------
+ # Latent prep
+ # ------------------------------------------------------------------
+
+ def prepare_latents(
+ self,
+ batch_size: int,
+ num_channels_latents: int,
+ height: int,
+ width: int,
+ dtype: torch.dtype,
+ device: torch.device,
+ generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
+ latents: Optional[torch.Tensor] = None,
+ ) -> torch.Tensor:
+ latent_h = height // self.vae_scale_factor
+ latent_w = width // self.vae_scale_factor
+ shape = (batch_size, latent_h * latent_w, num_channels_latents)
+ if latents is not None:
+ return latents.to(device=device, dtype=dtype)
+ return randn_tensor(shape, generator=generator, device=device, dtype=dtype)
+
+ # ------------------------------------------------------------------
+ # Input checks
+ # ------------------------------------------------------------------
+
+ def check_inputs(
+ self,
+ prompt,
+ height,
+ width,
+ prompt_embeds,
+ callback_on_step_end_tensor_inputs,
+ ) -> None:
+ if height is None or width is None:
+ raise ValueError(
+ "height and width must be provided (or use base_resolution + aspect_ratio)."
+ )
+ if height % self.vae_scale_factor or width % self.vae_scale_factor:
+ raise ValueError(
+ f"height and width must be divisible by {self.vae_scale_factor}; "
+ f"got ({height}, {width})."
+ )
+ if prompt is None and prompt_embeds is None:
+ raise ValueError("Either `prompt` or `prompt_embeds` must be provided.")
+ if callback_on_step_end_tensor_inputs is not None:
+ for k in callback_on_step_end_tensor_inputs:
+ if k not in self._callback_tensor_inputs:
+ raise ValueError(
+ f"callback_on_step_end_tensor_inputs entry {k!r} is not "
+ f"in {self._callback_tensor_inputs}."
+ )
+
+ # ------------------------------------------------------------------
+ # Decode
+ # ------------------------------------------------------------------
+
+ @staticmethod
+ def _patchify_latents(latents: torch.Tensor) -> torch.Tensor:
+ b, c, h, w = latents.shape
+ latents = latents.view(b, c, h // 2, 2, w // 2, 2)
+ latents = latents.permute(0, 1, 3, 5, 2, 4)
+ return latents.reshape(b, c * 4, h // 2, w // 2)
+
+ @staticmethod
+ def _unpatchify_latents(latents: torch.Tensor) -> torch.Tensor:
+ b, c, h, w = latents.shape
+ latents = latents.reshape(b, c // 4, 2, 2, h, w)
+ latents = latents.permute(0, 1, 4, 2, 5, 3)
+ return latents.reshape(b, c // 4, h * 2, w * 2)
+
+ @torch.no_grad()
+ def _decode(self, latents: torch.Tensor, latent_h: int, latent_w: int):
+ latents = rearrange(
+ latents,
+ "b (h w) (c p1 p2) -> b c (h p1) (w p2)",
+ p1=2, p2=2, h=latent_h, w=latent_w,
+ )
+ latents = latents.to(self.vae.dtype)
+ # Reverse the VAE latent normalization used by Lens. We compute the
+ # shift/scale at runtime from the live ``vae.bn`` so this stays correct
+ # under cpu-offload (where the VAE may be moved between devices).
+ bn = self.vae.bn
+ mean = bn.running_mean.view(1, -1, 1, 1)
+ var = bn.running_var.view(1, -1, 1, 1)
+ std = torch.sqrt(var + self.vae.config.batch_norm_eps)
+ shift = (-mean).to(device=latents.device, dtype=latents.dtype)
+ scale = (1.0 / std).to(device=latents.device, dtype=latents.dtype)
+ x = self._patchify_latents(latents)
+ x = x / scale - shift
+ x = self._unpatchify_latents(x)
+ return self.vae.decode(x).sample
+
+ @staticmethod
+ def _to_pil(image: torch.Tensor) -> List[Image.Image]:
+ # image: [B, C, H, W] in [-1, 1].
+ image = image.clamp(-1.0, 1.0)
+ image = (image + 1.0) * (255.0 / 2.0)
+ image = image.permute(0, 2, 3, 1).to(device="cpu", dtype=torch.uint8).numpy()
+ return [Image.fromarray(im) for im in image]
+
+ # ------------------------------------------------------------------
+ # __call__
+ # ------------------------------------------------------------------
+
+ @torch.no_grad()
+ def __call__(
+ self,
+ prompt: Union[str, List[str]] = None, # noqa: RUF013
+ negative_prompt: Union[str, List[str]] = "",
+ height: Optional[int] = None,
+ width: Optional[int] = None,
+ base_resolution: Optional[int] = None,
+ aspect_ratio: Optional[str] = None,
+ num_inference_steps: int = 50,
+ guidance_scale: float = 4.0,
+ num_images_per_prompt: int = 1,
+ generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
+ latents: Optional[torch.Tensor] = None,
+ prompt_embeds: Optional[List[torch.Tensor]] = None,
+ prompt_mask: Optional[torch.Tensor] = None,
+ negative_prompt_embeds: Optional[List[torch.Tensor]] = None,
+ negative_prompt_mask: Optional[torch.Tensor] = None,
+ output_type: str = "pil",
+ return_dict: bool = True,
+ callback_on_step_end: Optional[Callable[[Any, int, int, Dict], Dict]] = None,
+ callback_on_step_end_tensor_inputs: List[str] = ["latents"],
+ max_sequence_length: int = 512,
+ enable_reasoner: bool = False,
+ ):
+ # 0. Resolution defaulting.
+ if base_resolution is not None and aspect_ratio is not None:
+ height, width = resolve_resolution(base_resolution, aspect_ratio)
+ elif height is None or width is None:
+ height = width = self.default_sample_size
+
+ # 1. Input validation.
+ self.check_inputs(
+ prompt, height, width, prompt_embeds, callback_on_step_end_tensor_inputs
+ )
+
+ device = self._execution_device
+ dtype = self.transformer.dtype
+
+ # 2. Reasoner refinement (no-op when disabled and no API).
+ if prompt is not None:
+ prompts = [prompt] if isinstance(prompt, str) else list(prompt)
+ prompts = self.refine_prompt(prompts, enable_reasoner=enable_reasoner)
+ self._last_refined_prompts = prompts # pylint: disable=attribute-defined-outside-init
+ else:
+ prompts = None
+
+ # 3. Encode positives and negatives.
+ prompt_embeds, prompt_mask, negative_prompt_embeds, negative_prompt_mask = self.encode_prompt(
+ prompt=prompts,
+ negative_prompt=negative_prompt,
+ num_images_per_prompt=num_images_per_prompt,
+ prompt_embeds=prompt_embeds,
+ prompt_mask=prompt_mask,
+ negative_prompt_embeds=negative_prompt_embeds,
+ negative_prompt_mask=negative_prompt_mask,
+ max_sequence_length=max_sequence_length,
+ device=device,
+ )
+
+ # 4. Pad pos/neg to a shared S_txt for joint CFG batching.
+ prompt_embeds, prompt_mask, negative_prompt_embeds, negative_prompt_mask = self._align_text_features(
+ prompt_embeds, prompt_mask, negative_prompt_embeds, negative_prompt_mask
+ )
+
+ encoder_features = [
+ torch.cat([pf, nf], dim=0).to(dtype=dtype)
+ for pf, nf in zip(prompt_embeds, negative_prompt_embeds)
+ ]
+ encoder_mask = torch.cat([prompt_mask, negative_prompt_mask], dim=0)
+
+ # 5. Prepare latents.
+ batch_size = prompt_embeds[0].shape[0]
+ latent_h = height // self.vae_scale_factor
+ latent_w = width // self.vae_scale_factor
+ seq_len = latent_h * latent_w
+ latents = self.prepare_latents(
+ batch_size, self.latent_channels, height, width,
+ dtype=dtype, device=device, generator=generator, latents=latents,
+ )
+
+ # 6. Scheduler.
+ mu = compute_empirical_mu(seq_len, num_inference_steps)
+ sigmas = np.linspace(1.0, 1.0 / num_inference_steps, num_inference_steps)
+ self.scheduler.set_timesteps(sigmas=sigmas, device=device, mu=mu)
+
+ # 7. Denoising loop.
+ img_shapes = [(1, latent_h, latent_w)]
+ with self.progress_bar(total=num_inference_steps) as progress_bar:
+ for i, t in enumerate(self.scheduler.timesteps):
+ timestep = t.expand(batch_size * 2).to(latents.dtype)
+ hidden_states = latents.repeat(2, 1, 1)
+
+ noise = self.transformer(
+ hidden_states=hidden_states,
+ encoder_hidden_states=encoder_features,
+ encoder_hidden_states_mask=encoder_mask,
+ timestep=timestep / 1000,
+ img_shapes=img_shapes,
+ )
+
+ cond, uncond = noise.chunk(2)
+ comb = uncond + guidance_scale * (cond - uncond)
+ cond_norm = torch.norm(cond, dim=-1, keepdim=True)
+ comb_norm = torch.norm(comb, dim=-1, keepdim=True)
+ scale = torch.where(
+ comb_norm > 0,
+ cond_norm / comb_norm.clamp_min(1e-12),
+ torch.ones_like(comb_norm),
+ )
+ noise_pred = comb * scale
+
+ latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
+
+ if callback_on_step_end is not None:
+ cb_kwargs = {
+ k: locals()[k] for k in callback_on_step_end_tensor_inputs
+ }
+ cb_out = callback_on_step_end(self, i, t, cb_kwargs)
+ latents = cb_out.pop("latents", latents)
+ prompt_embeds = cb_out.pop("prompt_embeds", prompt_embeds)
+ negative_prompt_embeds = cb_out.pop(
+ "negative_prompt_embeds", negative_prompt_embeds
+ )
+
+ progress_bar.update()
+
+ # 8. Decode.
+ if output_type == "latent":
+ images: Any = latents
+ else:
+ decoded = self._decode(latents, latent_h, latent_w)
+ if output_type == "pil":
+ images = self._to_pil(decoded)
+ elif output_type == "np":
+ decoded = decoded.clamp(-1.0, 1.0)
+ decoded = (decoded + 1.0) * 0.5
+ images = decoded.permute(0, 2, 3, 1).to("cpu", torch.float32).numpy()
+ else:
+ raise ValueError(
+ f"output_type must be one of 'pil', 'np', 'latent'; got {output_type!r}."
+ )
+
+ self.maybe_free_model_hooks()
+
+ if not return_dict:
+ return (images,)
+ return LensPipelineOutput(images=images)
+
+ # ------------------------------------------------------------------
+ # Misc helpers
+ # ------------------------------------------------------------------
+
+ @staticmethod
+ def _align_text_features(
+ pos_features: List[torch.Tensor],
+ pos_mask: torch.Tensor,
+ neg_features: List[torch.Tensor],
+ neg_mask: torch.Tensor,
+ ):
+ """Pad pos/neg encodings and masks to a common ``S_txt``."""
+ if not pos_features or not neg_features:
+ raise ValueError("Positive and negative text feature lists must be non-empty.")
+ if len(pos_features) != len(neg_features):
+ raise ValueError(
+ "Positive and negative text feature lists must have the same "
+ f"number of layers; got {len(pos_features)} and {len(neg_features)}."
+ )
+ seq_pos = pos_features[0].shape[1]
+ seq_neg = neg_features[0].shape[1]
+ if pos_mask.shape[1] != seq_pos:
+ raise ValueError(
+ f"prompt_mask length {pos_mask.shape[1]} does not match "
+ f"prompt feature length {seq_pos}."
+ )
+ if pos_mask.shape[0] != pos_features[0].shape[0]:
+ raise ValueError(
+ f"prompt_mask batch size {pos_mask.shape[0]} does not match "
+ f"prompt feature batch size {pos_features[0].shape[0]}."
+ )
+ if neg_mask.shape[1] != seq_neg:
+ raise ValueError(
+ f"negative_prompt_mask length {neg_mask.shape[1]} does not "
+ f"match negative prompt feature length {seq_neg}."
+ )
+ if neg_mask.shape[0] != neg_features[0].shape[0]:
+ raise ValueError(
+ f"negative_prompt_mask batch size {neg_mask.shape[0]} does "
+ f"not match negative prompt feature batch size {neg_features[0].shape[0]}."
+ )
+ if pos_features[0].shape[0] != neg_features[0].shape[0]:
+ raise ValueError(
+ "Positive and negative text features must have the same batch "
+ f"size; got {pos_features[0].shape[0]} and {neg_features[0].shape[0]}."
+ )
+ for i, feat in enumerate(pos_features):
+ if feat.shape[:2] != pos_features[0].shape[:2]:
+ raise ValueError(
+ f"Positive feature layer {i} shape {feat.shape[:2]} does "
+ f"not match layer 0 shape {pos_features[0].shape[:2]}."
+ )
+ for i, feat in enumerate(neg_features):
+ if feat.shape[:2] != neg_features[0].shape[:2]:
+ raise ValueError(
+ f"Negative feature layer {i} shape {feat.shape[:2]} does "
+ f"not match layer 0 shape {neg_features[0].shape[:2]}."
+ )
+
+ target = max(seq_pos, seq_neg)
+
+ def pad(features: List[torch.Tensor], cur: int) -> List[torch.Tensor]:
+ if cur == target:
+ return features
+ pad_len = target - cur
+ return [
+ torch.cat(
+ [feat, feat.new_zeros((feat.shape[0], pad_len, feat.shape[-1]))],
+ dim=1,
+ )
+ for feat in features
+ ]
+
+ def pad_mask(mask: torch.Tensor, cur: int) -> torch.Tensor:
+ if cur == target:
+ return mask
+ return torch.cat(
+ [
+ mask,
+ torch.zeros(
+ (mask.shape[0], target - cur),
+ dtype=torch.bool, device=mask.device,
+ ),
+ ],
+ dim=1,
+ )
+
+ pos_features = pad(pos_features, seq_pos)
+ neg_features = pad(neg_features, seq_neg)
+ pos_mask = pad_mask(pos_mask.bool(), seq_pos)
+ neg_mask = pad_mask(neg_mask.bool(), seq_neg)
+ return pos_features, pos_mask, neg_features, neg_mask
diff --git a/pipelines/lens/reasoner.py b/pipelines/lens/reasoner.py
new file mode 100644
index 000000000..7814ed7f5
--- /dev/null
+++ b/pipelines/lens/reasoner.py
@@ -0,0 +1,252 @@
+"""Prompt reasoner - refines user prompts before they hit the text encoder.
+
+Decision matrix (driven by ``enable`` and whether an OpenAI-compatible API
+is configured):
+
+| ``enable`` | OpenAI API set? | Behavior |
+| ----------------- | --------------- | ----------------------------------------- |
+| ``False`` (default) | no | identity (return prompts unchanged) |
+| ``False`` | yes | refine via OpenAI-compatible API |
+| ``True`` | no | refine via the local GPT-OSS |
+| ``True`` | yes | refine via OpenAI-compatible API |
+
+The OpenAI path uses any chat-completion endpoint speaking the OpenAI v1
+schema (e.g. ``vllm``, ``ollama --openai-compat``, ``together.ai``).
+"""
+
+from __future__ import annotations
+
+import re
+from typing import List, Optional, Sequence
+
+import torch
+
+
+THINK_BLOCK_RE = re.compile(r".*?", re.DOTALL | re.IGNORECASE)
+HARMONY_FINAL_RE = re.compile(
+ r"<\|start\|>assistant(?:<\|channel\|>final)?<\|message\|>(.*?)(?:<\|return\|>|<\|end\|>|$)",
+ re.DOTALL,
+)
+HARMONY_DIRECT_FINAL_RE = re.compile(
+ r"<\|channel\|>final<\|message\|>(.*?)(?:<\|return\|>|<\|end\|>|$)",
+ re.DOTALL,
+)
+PLAIN_HARMONY_FINAL_MARKER_RE = re.compile(r"assistant\s*final\s*", re.IGNORECASE)
+PLAIN_HARMONY_DIRECT_FINAL_RE = re.compile(r"(?:^|\n)\s*final\s*", re.IGNORECASE)
+
+
+SYSTEM_PROMPT = """
+You are a prompt rewriter for a text-to-image model.
+Your task is to convert the user's input into a single, precise, descriptive image prompt suitable for a text-to-image model.
+Follow these rules strictly:
+
+1. The output must be a clear and accurate description of a single image scene, written in the style of a text-to-image prompt.
+ - Do not include explanations, reasoning, commentary, or meta text.
+ - Do not ask questions.
+ - Do not output multiple options.
+ - Do not use uncertain, speculative, or alternative wording such as "maybe", "possibly", "perhaps", "or", "might", or "could".
+
+2. Preserve the user's intended scene faithfully.
+ - Do not change the objects, entities, attributes, actions, relationships, or core setting explicitly described by the user.
+ - You may add reasonable visual details only when they help make the image concrete and coherent.
+ - Any added details must be consistent with the user's description and must not introduce new important objects or alter the meaning.
+
+3. If the image contains many main subjects of the same kind, describe each subject in detail, including humans, animals, objects, and any other prominent elements.
+ - For each subject, include its appearance, color, size, shape, material, pose, expression, and position if applicable in the scene.
+ - Make sure every main subject is clearly distinguishable from the others, such as in a scene with "4 dogs," describing each dog separately.
+
+4. The output must fully cover the scene implied by the user's input.
+ - Include the main subjects, relevant attributes, actions, spatial relationships, environment, and visible details necessary to render the scene.
+ - If the user input is already sufficiently detailed and already suitable for image generation, keep it unchanged or only make minimal edits for fluency and clarity.
+
+5. Resolve content that requires simple inference into explicit visual results when the result is unambiguous and visually representable.
+ - Example: if the user says "the answer to 2+2 is written on the blackboard", output should explicitly describe "the blackboard shows 2+2=4".
+ - Use only direct, necessary inference that is clearly implied by the user input.
+ - Do not invent hidden facts, backstory, or ambiguous details.
+
+6. Language rule:
+ - If the user input is not in English, output in the same language.
+ - Otherwise, output in English.
+
+7. Output format:
+ - Output exactly one final rewritten prompt.
+ - Do not use bullet points, numbering, JSON, XML, Markdown, or quotation marks unless they are part of the scene itself.
+
+Your goal is to produce a prompt that is concrete, visual, faithful to the user intent, and directly usable as input to a text-to-image model.
+""".strip()
+
+
+def _extract_plain_harmony_final(text: str) -> Optional[str]:
+ matches = list(PLAIN_HARMONY_FINAL_MARKER_RE.finditer(text))
+ if matches:
+ final_text = text[matches[-1].end() :].strip()
+ return final_text or None
+
+ if text.lstrip().lower().startswith("analysis"):
+ matches = list(PLAIN_HARMONY_DIRECT_FINAL_RE.finditer(text))
+ if matches:
+ final_text = text[matches[-1].end() :].strip()
+ return final_text or None
+ return None
+
+
+def _clean_reasoner_output(text: str) -> str:
+ text = text.strip()
+ final_match = None
+ for match in HARMONY_FINAL_RE.finditer(text):
+ final_match = match
+ if final_match is not None:
+ text = final_match.group(1).strip()
+ else:
+ direct_final_match = None
+ for match in HARMONY_DIRECT_FINAL_RE.finditer(text):
+ direct_final_match = match
+ if direct_final_match is not None:
+ text = direct_final_match.group(1).strip()
+ else:
+ plain_final = _extract_plain_harmony_final(text)
+ if plain_final is not None:
+ text = plain_final
+
+ text = THINK_BLOCK_RE.sub("", text).strip()
+ if "" in text.lower():
+ text = re.split(r"", text, flags=re.IGNORECASE)[-1].strip()
+ plain_final = _extract_plain_harmony_final(text)
+ if plain_final is not None:
+ text = plain_final
+ for token in (
+ "<|channel|>analysis<|message|>",
+ "<|start|>assistant<|channel|>analysis<|message|>",
+ "<|channel|>final<|message|>",
+ "<|start|>assistant<|channel|>final<|message|>",
+ "<|start|>assistant<|message|>",
+ "<|return|>",
+ "<|end|>",
+ "<|endoftext|>",
+ "<|im_end|>",
+ ):
+ text = text.replace(token, "")
+
+ text = text.strip()
+ if re.match(r"^(?:analysis|assistant\s*analysis)(?:\b|[A-Z])", text, flags=re.IGNORECASE | re.DOTALL):
+ return ""
+ if text.startswith("```") and text.endswith("```"):
+ lines = text.splitlines()
+ if len(lines) >= 3:
+ text = "\n".join(lines[1:-1]).strip()
+ if len(text) >= 2 and text[0] == text[-1] == '"':
+ text = text[1:-1].strip()
+ return " ".join(text.split())
+
+
+class PromptReasoner:
+ """Optional prompt rewriter, used by ``LensPipeline.refine_prompt``."""
+
+ def __init__(
+ self,
+ *,
+ text_encoder=None,
+ tokenizer=None,
+ openai_api_key: Optional[str] = None,
+ openai_base_url: Optional[str] = None,
+ openai_model: Optional[str] = None,
+ max_new_tokens: int = 4096,
+ temperature: float = 0.7,
+ ) -> None:
+ self.text_encoder = text_encoder
+ self.tokenizer = tokenizer
+ self.openai_api_key = openai_api_key
+ self.openai_base_url = openai_base_url
+ self.openai_model = openai_model
+ self.max_new_tokens = int(max_new_tokens)
+ self.temperature = float(temperature)
+ self._client = None # lazily constructed
+
+ @property
+ def has_api(self) -> bool:
+ return bool(self.openai_api_key and self.openai_model)
+
+ def refine(self, prompts: Sequence[str], enable: bool) -> List[str]:
+ prompts = list(prompts)
+ # API takes precedence whenever it is configured.
+ if self.has_api:
+ return self._refine_via_api(prompts)
+ if enable:
+ if self.text_encoder is None or self.tokenizer is None:
+ raise RuntimeError(
+ "Reasoner enabled with no API: both text_encoder and "
+ "tokenizer must be provided to use the local GPT-OSS as "
+ "the reasoner."
+ )
+ return self._refine_via_local(prompts)
+ return prompts
+
+ # ------------------------------------------------------------------
+ # Local GPT-OSS path
+ # ------------------------------------------------------------------
+
+ @torch.no_grad()
+ def _refine_via_local(self, prompts: List[str]) -> List[str]:
+ refined: List[str] = []
+ for prompt in prompts:
+ system_prompt = (
+ f"{SYSTEM_PROMPT}\n\n"
+ "Keep any reasoning private. The visible answer must contain only the final rewritten prompt."
+ )
+ conversation = [
+ {"role": "system", "content": system_prompt, "thinking": None},
+ {"role": "user", "content": prompt, "thinking": None},
+ ]
+ text = self.tokenizer.apply_chat_template(
+ conversation, tokenize=False, add_generation_prompt=True, reasoning_effort="low"
+ )
+ input_ids = self.tokenizer(
+ text, return_tensors="pt", add_special_tokens=True
+ ).input_ids
+ out_ids = self.text_encoder.generate(
+ input_ids,
+ max_new_tokens=self.max_new_tokens,
+ do_sample=self.temperature > 0.0,
+ temperature=max(self.temperature, 1e-5),
+ pad_token_id=self.tokenizer.pad_token_id,
+ )
+ new_tokens = out_ids[0, input_ids.shape[1]:]
+ text_out = self.tokenizer.decode(new_tokens, skip_special_tokens=False)
+ clean_text_out = _clean_reasoner_output(text_out)
+ refined.append(clean_text_out or prompt)
+ return refined
+
+ # ------------------------------------------------------------------
+ # OpenAI-compatible API path
+ # ------------------------------------------------------------------
+
+ def _client_or_raise(self):
+ if self._client is None:
+ try:
+ from openai import OpenAI
+ except ImportError as exc:
+ raise ImportError(
+ "openai package not installed. `pip install openai` to use "
+ "the API-based reasoner."
+ ) from exc
+ self._client = OpenAI(
+ api_key=self.openai_api_key,
+ base_url=self.openai_base_url,
+ )
+ return self._client
+
+ def _refine_via_api(self, prompts: List[str]) -> List[str]:
+ client = self._client_or_raise()
+ out: List[str] = []
+ for prompt in prompts:
+ resp = client.chat.completions.create(
+ model=self.openai_model,
+ messages=[
+ {"role": "system", "content": SYSTEM_PROMPT},
+ {"role": "user", "content": prompt},
+ ],
+ max_tokens=self.max_new_tokens,
+ )
+ text = (resp.choices[0].message.content or "").strip()
+ out.append(text or prompt)
+ return out
diff --git a/pipelines/lens/resolution.py b/pipelines/lens/resolution.py
new file mode 100644
index 000000000..1dae67eda
--- /dev/null
+++ b/pipelines/lens/resolution.py
@@ -0,0 +1,62 @@
+"""Resolution buckets for Lens inference.
+
+Two base resolutions (1024 and 1440) crossed with nine aspect ratios. All
+heights/widths are divisible by 16 so they tile cleanly into Flux2 VAE latents
+(downsample factor 16).
+"""
+
+from __future__ import annotations
+
+from typing import Dict, Tuple
+
+
+# Bucket maps. Keys are aspect ratios as "W:H" strings; values are (height, width).
+# 1440 base buckets are computed as round_to_16(1024_value * 1440 / 1024).
+RESOLUTION_BUCKETS: Dict[int, Dict[str, Tuple[int, int]]] = {
+ 1024: {
+ "1:2": (1472, 736),
+ "9:16": (1376, 768),
+ "2:3": (1248, 832),
+ "3:4": (1152, 864),
+ "1:1": (1024, 1024),
+ "4:3": ( 864, 1152),
+ "3:2": ( 832, 1248),
+ "16:9": ( 768, 1376),
+ "2:1": ( 736, 1472),
+ },
+ 1440: {
+ "1:2": (2080, 1040),
+ "9:16": (1936, 1088),
+ "2:3": (1760, 1168),
+ "3:4": (1616, 1216),
+ "1:1": (1440, 1440),
+ "4:3": (1216, 1616),
+ "3:2": (1168, 1760),
+ "16:9": (1088, 1936),
+ "2:1": (1040, 2080),
+ },
+}
+
+
+SUPPORTED_BASE_RESOLUTIONS = tuple(RESOLUTION_BUCKETS.keys())
+SUPPORTED_ASPECT_RATIOS = tuple(RESOLUTION_BUCKETS[1024].keys())
+
+
+def resolve_resolution(base_resolution: int, aspect_ratio: str) -> Tuple[int, int]:
+ """Return (height, width) for the requested bucket.
+
+ Aspect ratio is interpreted as W:H (e.g. "16:9" is landscape,
+ "9:16" is portrait).
+ """
+ if base_resolution not in RESOLUTION_BUCKETS:
+ raise ValueError(
+ f"Unsupported base_resolution={base_resolution}. "
+ f"Supported: {SUPPORTED_BASE_RESOLUTIONS}"
+ )
+ table = RESOLUTION_BUCKETS[base_resolution]
+ if aspect_ratio not in table:
+ raise ValueError(
+ f"Unsupported aspect_ratio={aspect_ratio!r}. "
+ f"Supported: {SUPPORTED_ASPECT_RATIOS}"
+ )
+ return table[aspect_ratio]
diff --git a/pipelines/lens/text_encoder.py b/pipelines/lens/text_encoder.py
new file mode 100644
index 000000000..0b16410e4
--- /dev/null
+++ b/pipelines/lens/text_encoder.py
@@ -0,0 +1,137 @@
+"""GPT-OSS text encoder for Lens.
+
+We subclass ``transformers.GptOssForCausalLM`` so we can:
+
+1. Return hidden states *only* at a configured layer subset (default
+ ``[5, 11, 17, 23]``), avoiding the memory cost of HF's stock
+ ``output_hidden_states=True`` which materializes every layer.
+2. Early-exit after the last selected layer, since we don't need the
+ downstream LM head at all when extracting features.
+
+Standard ``generate(...)`` is inherited unchanged and is used by the optional
+prompt reasoner.
+"""
+
+from __future__ import annotations
+
+from typing import List, Optional, Sequence
+
+import torch
+from transformers.masking_utils import (
+ create_causal_mask,
+ create_sliding_window_causal_mask,
+)
+from transformers.models.gpt_oss.modeling_gpt_oss import GptOssForCausalLM
+
+
+class LensGptOssEncoder(GptOssForCausalLM):
+ """``GptOssForCausalLM`` subclass that exposes selected hidden states."""
+
+ def set_selected_layers(self, layer_indices: Sequence[int]) -> None:
+ layers = [int(i) for i in layer_indices]
+ if not layers:
+ raise ValueError("layer_indices must be non-empty")
+ if len(set(layers)) != len(layers):
+ raise ValueError(f"layer_indices must be unique; got {layers}")
+ if min(layers) < 0 or max(layers) >= len(self.model.layers):
+ raise ValueError(
+ f"layer_indices out of range; got {layers}, "
+ f"model has {len(self.model.layers)} layers"
+ )
+ self._lens_selected_layers = layers
+ self._lens_max_layer = max(layers)
+
+ @torch.no_grad()
+ def forward( # type: ignore[override]
+ self,
+ input_ids: Optional[torch.LongTensor] = None,
+ attention_mask: Optional[torch.Tensor] = None,
+ *args,
+ **kwargs,
+ ):
+ """Lens-specific forward.
+
+ When ``input_ids`` and ``attention_mask`` are provided AND
+ ``set_selected_layers(...)`` has been called, this returns the list of
+ hidden states at the configured selected layers (the Lens feature
+ extraction path).
+
+ Otherwise, falls back to ``GptOssForCausalLM.forward`` so that
+ ``generate(...)`` (used by the prompt reasoner) still works unchanged.
+ """
+ is_lens_feature_call = (
+ input_ids is not None
+ and attention_mask is not None
+ and hasattr(self, "_lens_selected_layers")
+ and not args
+ and not kwargs
+ )
+
+ target_device = self.model.embed_tokens.weight.device
+ if input_ids is not None and input_ids.device != target_device:
+ input_ids = input_ids.to(target_device)
+ if attention_mask is not None and attention_mask.device != target_device:
+ attention_mask = attention_mask.to(target_device)
+
+ if not is_lens_feature_call:
+ return super().forward(input_ids, attention_mask, *args, **kwargs)
+
+ model = self.model
+ inputs_embeds = model.embed_tokens(input_ids)
+ position_ids = torch.arange(
+ inputs_embeds.shape[1], device=inputs_embeds.device
+ ).unsqueeze(0).expand_as(input_ids)
+
+ mask_kwargs = {
+ "config": model.config,
+ "inputs_embeds": inputs_embeds,
+ "attention_mask": attention_mask,
+ "past_key_values": None,
+ "position_ids": position_ids,
+ }
+ causal_mask_mapping = {
+ "full_attention": create_causal_mask(**mask_kwargs),
+ "sliding_attention": create_sliding_window_causal_mask(**mask_kwargs),
+ }
+
+ hidden_states = inputs_embeds
+ position_embeddings = model.rotary_emb(hidden_states, position_ids)
+
+ captured: List[torch.Tensor] = [None] * len(self._lens_selected_layers)
+ index_lookup = {idx: pos for pos, idx in enumerate(self._lens_selected_layers)}
+
+ for i, decoder_layer in enumerate(model.layers):
+ hidden_states = decoder_layer(
+ hidden_states,
+ attention_mask=causal_mask_mapping[model.config.layer_types[i]],
+ position_embeddings=position_embeddings,
+ position_ids=position_ids,
+ past_key_values=None,
+ use_cache=False,
+ )
+ if i in index_lookup:
+ captured[index_lookup[i]] = hidden_states
+ if i == self._lens_max_layer:
+ break
+
+ for pos, layer_idx in enumerate(self._lens_selected_layers):
+ if captured[pos] is None:
+ raise RuntimeError(
+ f"Failed to capture hidden state for layer {layer_idx}"
+ )
+ return captured
+
+ def encode_layers(
+ self,
+ input_ids: torch.LongTensor,
+ attention_mask: torch.Tensor,
+ ) -> List[torch.Tensor]:
+ """Backwards-compatible alias for the Lens feature path.
+
+ Kept so existing call sites (``LensPipeline._get_text_embeddings``,
+ external users) keep working. New code should call the encoder
+ directly: ``encoder(input_ids, attention_mask)``.
+ """
+ if not hasattr(self, "_lens_selected_layers"):
+ raise RuntimeError("Call set_selected_layers(...) before encode_layers().")
+ return self(input_ids=input_ids, attention_mask=attention_mask)
diff --git a/pipelines/lens/transformer.py b/pipelines/lens/transformer.py
new file mode 100644
index 000000000..799b91929
--- /dev/null
+++ b/pipelines/lens/transformer.py
@@ -0,0 +1,554 @@
+"""Lens denoising transformer (DiT).
+
+The model uses a double-stream architecture with joint image+text attention,
+RoPE on both streams, and SwiGLU MLPs.
+"""
+
+from __future__ import annotations
+
+import math
+from typing import Any, Dict, List, Optional, Tuple, Union
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from diffusers.configuration_utils import ConfigMixin, register_to_config
+from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
+from diffusers.models.attention import FeedForward
+from diffusers.models.cache_utils import CacheMixin
+from diffusers.models.embeddings import TimestepEmbedding, Timesteps
+from diffusers.models.modeling_utils import ModelMixin
+from diffusers.models.normalization import AdaLayerNormContinuous, RMSNorm
+
+
+# ---------------------------------------------------------------------------
+# Embeddings & RoPE
+# ---------------------------------------------------------------------------
+
+
+def get_timestep_embedding(
+ timesteps: torch.Tensor,
+ embedding_dim: int,
+ flip_sin_to_cos: bool = False,
+ downscale_freq_shift: float = 1.0,
+ scale: float = 1.0,
+ max_period: int = 10000,
+) -> torch.Tensor:
+ """Sinusoidal timestep embeddings (DDPM-style)."""
+ assert timesteps.ndim == 1, "Timesteps should be 1-D"
+ half_dim = embedding_dim // 2
+ exponent = -math.log(max_period) * torch.arange(
+ 0, half_dim, dtype=torch.float32, device=timesteps.device
+ )
+ exponent = exponent / (half_dim - downscale_freq_shift)
+ emb = torch.exp(exponent).to(timesteps.dtype)
+ emb = timesteps[:, None].float() * emb[None, :]
+ emb = scale * emb
+ emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
+ if flip_sin_to_cos:
+ emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
+ if embedding_dim % 2 == 1:
+ emb = F.pad(emb, (0, 1, 0, 0))
+ return emb
+
+
+def apply_rotary_emb_lens(
+ x: torch.Tensor,
+ freqs_cis: torch.Tensor,
+) -> torch.Tensor:
+ """Apply complex-valued RoPE (Lens variant).
+
+ Args:
+ x: [B, S, H, D] query or key tensor.
+ freqs_cis: [S, D/2] complex tensor of rotation factors.
+ """
+ x_complex = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
+ freqs_cis = freqs_cis.unsqueeze(1) # broadcast over heads
+ x_out = torch.view_as_real(x_complex * freqs_cis).flatten(3)
+ return x_out.type_as(x)
+
+
+class GateMLP(nn.Module):
+ """SwiGLU MLP used by the transformer blocks."""
+
+ def __init__(self, dim: int, hidden_dim: int) -> None:
+ super().__init__()
+ self.w1 = nn.Linear(dim, hidden_dim, bias=False)
+ self.w2 = nn.Linear(hidden_dim, dim, bias=False)
+ self.w3 = nn.Linear(dim, hidden_dim, bias=False)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ return self.w2(F.silu(self.w1(x)) * self.w3(x))
+
+
+class LensTimestepProjEmbeddings(nn.Module):
+ def __init__(self, embedding_dim: int) -> None:
+ super().__init__()
+ self.time_proj = Timesteps(
+ num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0, scale=1000
+ )
+ self.timestep_embedder = TimestepEmbedding(
+ in_channels=256, time_embed_dim=embedding_dim
+ )
+
+ def forward(self, timestep: torch.Tensor, hidden_states: torch.Tensor) -> torch.Tensor:
+ proj = self.time_proj(timestep)
+ return self.timestep_embedder(proj.to(dtype=hidden_states.dtype))
+
+
+class LensEmbedRope(nn.Module):
+ """Frame/H/W axial RoPE shared between image and text streams."""
+
+ def __init__(self, theta: int, axes_dim: List[int], scale_rope: bool = False) -> None:
+ super().__init__()
+ self.theta = theta
+ self.axes_dim = axes_dim
+ self.scale_rope = scale_rope
+ pos_index = torch.arange(4096)
+ neg_index = torch.arange(4096).flip(0) * -1 - 1
+ self.pos_freqs = torch.cat(
+ [self._rope_params(pos_index, d, theta) for d in axes_dim], dim=1
+ )
+ self.neg_freqs = torch.cat(
+ [self._rope_params(neg_index, d, theta) for d in axes_dim], dim=1
+ )
+ # Note: we deliberately do NOT register these as buffers - registering
+ # complex tensors as buffers strips the imaginary component on save/load.
+ self.rope_cache: Dict[str, torch.Tensor] = {}
+
+ @staticmethod
+ def _rope_params(index: torch.Tensor, dim: int, theta: int = 10000) -> torch.Tensor:
+ assert dim % 2 == 0
+ freqs = torch.outer(
+ index, 1.0 / torch.pow(theta, torch.arange(0, dim, 2).float().div(dim))
+ )
+ return torch.polar(torch.ones_like(freqs), freqs)
+
+ def forward(
+ self,
+ video_fhw: Union[List[Tuple[int, int, int]], Tuple[int, int, int]],
+ txt_seq_lens: Union[List[int], int],
+ device: torch.device = torch.device("cuda"),
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ if self.pos_freqs.device != device:
+ self.pos_freqs = self.pos_freqs.to(device)
+ self.neg_freqs = self.neg_freqs.to(device)
+
+ if isinstance(video_fhw, list):
+ video_fhw = video_fhw[0]
+ if not isinstance(video_fhw, list):
+ video_fhw = [video_fhw]
+ if not isinstance(txt_seq_lens, list):
+ txt_seq_lens = [txt_seq_lens]
+ assert len(video_fhw) == 1, "video_fhw must have length 1"
+
+ vid_freqs = []
+ max_vid_index = 0
+ for idx, fhw in enumerate(video_fhw):
+ frame, height, width = fhw
+ rope_key = f"{idx}_{height}_{width}"
+ if rope_key not in self.rope_cache:
+ self.rope_cache[rope_key] = (
+ self._compute_video_freqs(frame, height, width, idx=0).to("cpu")
+ )
+ video_freq = self.rope_cache[rope_key].to(device)
+ if self.scale_rope:
+ max_vid_index = max(height // 2, width // 2, max_vid_index)
+ else:
+ max_vid_index = max(height, width, max_vid_index)
+ vid_freqs.append(video_freq)
+
+ max_len = max(txt_seq_lens)
+ txt_freqs = self.pos_freqs[max_vid_index : max_vid_index + max_len, ...]
+ return torch.cat(vid_freqs, dim=0), txt_freqs
+
+ def _compute_video_freqs(self, frame: int, height: int, width: int, idx: int = 0) -> torch.Tensor:
+ seq_lens = frame * height * width
+ freqs_pos = self.pos_freqs.split([d // 2 for d in self.axes_dim], dim=1)
+ freqs_neg = self.neg_freqs.split([d // 2 for d in self.axes_dim], dim=1)
+
+ freqs_frame = freqs_pos[0][idx : idx + frame].view(frame, 1, 1, -1).expand(frame, height, width, -1)
+ if self.scale_rope:
+ freqs_height = torch.cat(
+ [freqs_neg[1][-(height - height // 2) :], freqs_pos[1][: height // 2]], dim=0
+ ).view(1, height, 1, -1).expand(frame, height, width, -1)
+ freqs_width = torch.cat(
+ [freqs_neg[2][-(width - width // 2) :], freqs_pos[2][: width // 2]], dim=0
+ ).view(1, 1, width, -1).expand(frame, height, width, -1)
+ else:
+ freqs_height = freqs_pos[1][:height].view(1, height, 1, -1).expand(frame, height, width, -1)
+ freqs_width = freqs_pos[2][:width].view(1, 1, width, -1).expand(frame, height, width, -1)
+
+ freqs = torch.cat([freqs_frame, freqs_height, freqs_width], dim=-1).reshape(seq_lens, -1)
+ return freqs.clone().contiguous()
+
+
+# ---------------------------------------------------------------------------
+# Attention (joint image + text, plain SDPA)
+# ---------------------------------------------------------------------------
+
+
+class LensJointAttention(nn.Module):
+ """Joint image+text attention with fused QKV and SDPA backend."""
+
+ def __init__(
+ self,
+ query_dim: int,
+ added_kv_proj_dim: int,
+ dim_head: int = 64,
+ heads: int = 8,
+ out_dim: Optional[int] = None,
+ eps: float = 1e-5,
+ ) -> None:
+ super().__init__()
+ self.inner_dim = out_dim if out_dim is not None else dim_head * heads
+ self.heads = self.inner_dim // dim_head
+ self.dim_head = dim_head
+ self.out_dim = out_dim if out_dim is not None else query_dim
+
+ self.norm_q = RMSNorm(dim_head, eps=eps)
+ self.norm_k = RMSNorm(dim_head, eps=eps)
+ self.norm_added_q = RMSNorm(dim_head, eps=eps)
+ self.norm_added_k = RMSNorm(dim_head, eps=eps)
+
+ self.img_qkv = nn.Linear(query_dim, 3 * self.inner_dim, bias=True)
+ self.txt_qkv = nn.Linear(added_kv_proj_dim, 3 * self.inner_dim, bias=True)
+
+ self.to_out = nn.ModuleList([nn.Linear(self.inner_dim, self.out_dim, bias=True), nn.Identity()])
+ self.to_add_out = nn.Linear(self.inner_dim, query_dim, bias=True)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor,
+ image_rotary_emb: Tuple[torch.Tensor, torch.Tensor],
+ attention_mask: Optional[torch.Tensor] = None,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ bsz, seq_img, _ = hidden_states.shape
+ seq_txt = encoder_hidden_states.shape[1]
+
+ # Fused QKV per stream -> split.
+ img_qkv = self.img_qkv(hidden_states).view(bsz, seq_img, 3, self.heads, self.dim_head)
+ txt_qkv = self.txt_qkv(encoder_hidden_states).view(bsz, seq_txt, 3, self.heads, self.dim_head)
+ img_q, img_k, img_v = img_qkv.unbind(dim=2)
+ txt_q, txt_k, txt_v = txt_qkv.unbind(dim=2)
+
+ # QK RMSNorm.
+ img_q = self.norm_q(img_q)
+ img_k = self.norm_k(img_k)
+ txt_q = self.norm_added_q(txt_q)
+ txt_k = self.norm_added_k(txt_k)
+
+ # RoPE.
+ img_freqs, txt_freqs = image_rotary_emb
+ if img_freqs.shape[0] < seq_img:
+ raise ValueError(
+ f"Image RoPE length {img_freqs.shape[0]} is shorter than "
+ f"image sequence length {seq_img}."
+ )
+ img_freqs = img_freqs[:seq_img]
+ img_q = apply_rotary_emb_lens(img_q, img_freqs)
+ img_k = apply_rotary_emb_lens(img_k, img_freqs)
+ if seq_txt > 0:
+ if txt_freqs.shape[0] < seq_txt:
+ raise ValueError(
+ f"Text RoPE length {txt_freqs.shape[0]} is shorter than "
+ f"text sequence length {seq_txt}."
+ )
+ txt_freqs = txt_freqs[:seq_txt]
+ txt_q = apply_rotary_emb_lens(txt_q, txt_freqs)
+ txt_k = apply_rotary_emb_lens(txt_k, txt_freqs)
+
+ # Joint sequence per sample, then SDPA in [B, H, S, D] layout.
+ q = torch.cat([img_q, txt_q], dim=1).transpose(1, 2)
+ k = torch.cat([img_k, txt_k], dim=1).transpose(1, 2)
+ v = torch.cat([img_v, txt_v], dim=1).transpose(1, 2)
+
+ if attention_mask is not None:
+ expected_mask_shape = (bsz, 1, 1, seq_img + seq_txt)
+ if attention_mask.shape != expected_mask_shape:
+ raise ValueError(
+ f"attention_mask must have shape {expected_mask_shape}, "
+ f"got {tuple(attention_mask.shape)}."
+ )
+ attention_mask = attention_mask.to(q.dtype)
+ out = F.scaled_dot_product_attention(q, k, v, attn_mask=attention_mask)
+ out = out.transpose(1, 2).reshape(bsz, seq_img + seq_txt, -1)
+
+ img_out = self.to_out[1](self.to_out[0](out[:, :seq_img, :]))
+ txt_out = self.to_add_out(out[:, seq_img:, :])
+ return img_out, txt_out
+
+
+# ---------------------------------------------------------------------------
+# Transformer block
+# ---------------------------------------------------------------------------
+
+
+class LensTransformerBlock(nn.Module):
+ def __init__(
+ self,
+ dim: int,
+ num_attention_heads: int,
+ attention_head_dim: int,
+ eps: float = 1e-6,
+ rms_norm: bool = False,
+ gate_mlp: bool = False,
+ ) -> None:
+ super().__init__()
+ self.attn = LensJointAttention(
+ query_dim=dim,
+ added_kv_proj_dim=dim,
+ dim_head=attention_head_dim,
+ heads=num_attention_heads,
+ out_dim=dim,
+ eps=eps,
+ )
+
+ norm_cls = (lambda d: RMSNorm(d, eps=eps)) if rms_norm else (
+ lambda d: nn.LayerNorm(d, elementwise_affine=False, eps=eps)
+ )
+ if gate_mlp:
+ mlp_cls = lambda: GateMLP(dim, int(dim / 3 * 8)) # pylint: disable=unnecessary-lambda-assignment
+ else:
+ mlp_cls = lambda: FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") # pylint: disable=unnecessary-lambda-assignment
+
+ self.img_mod = nn.Sequential(nn.SiLU(), nn.Linear(dim, 6 * dim, bias=True))
+ self.img_norm1 = norm_cls(dim)
+ self.img_norm2 = norm_cls(dim)
+ self.img_mlp = mlp_cls()
+
+ self.txt_mod = nn.Sequential(nn.SiLU(), nn.Linear(dim, 6 * dim, bias=True))
+ self.txt_norm1 = norm_cls(dim)
+ self.txt_norm2 = norm_cls(dim)
+ self.txt_mlp = mlp_cls()
+
+ @staticmethod
+ def _modulate(x: torch.Tensor, mod_params: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
+ shift, scale, gate = mod_params.chunk(3, dim=-1)
+ return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1), gate.unsqueeze(1)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor,
+ temb: torch.Tensor,
+ image_rotary_emb: Tuple[torch.Tensor, torch.Tensor],
+ attention_mask: Optional[torch.Tensor] = None,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ img_mod1, img_mod2 = self.img_mod(temb).chunk(2, dim=-1)
+ txt_mod1, txt_mod2 = self.txt_mod(temb).chunk(2, dim=-1)
+
+ img_modulated, img_gate1 = self._modulate(self.img_norm1(hidden_states), img_mod1)
+ txt_modulated, txt_gate1 = self._modulate(self.txt_norm1(encoder_hidden_states), txt_mod1)
+
+ img_attn, txt_attn = self.attn(
+ hidden_states=img_modulated,
+ encoder_hidden_states=txt_modulated,
+ image_rotary_emb=image_rotary_emb,
+ attention_mask=attention_mask,
+ )
+
+ hidden_states = hidden_states + img_gate1 * img_attn
+ encoder_hidden_states = encoder_hidden_states + txt_gate1 * txt_attn
+
+ img_modulated2, img_gate2 = self._modulate(self.img_norm2(hidden_states), img_mod2)
+ hidden_states = hidden_states + img_gate2 * self.img_mlp(img_modulated2)
+
+ txt_modulated2, txt_gate2 = self._modulate(self.txt_norm2(encoder_hidden_states), txt_mod2)
+ encoder_hidden_states = encoder_hidden_states + txt_gate2 * self.txt_mlp(txt_modulated2)
+
+ return encoder_hidden_states, hidden_states
+
+
+# ---------------------------------------------------------------------------
+# Top-level model
+# ---------------------------------------------------------------------------
+
+
+class LensTransformer2DModel(
+ ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin
+):
+ """The Lens text-to-image DiT.
+
+ Supports a single conditioning stream of multi-layer text features. The
+ text features are normalized per layer, concatenated along the channel
+ axis, and projected to `inner_dim` before joining the image stream.
+ """
+
+ _supports_gradient_checkpointing = True
+ _no_split_modules = ["LensTransformerBlock"]
+ _skip_layerwise_casting_patterns = ["pos_embed", "norm"]
+ _repeated_blocks = ["LensTransformerBlock"]
+
+ @register_to_config
+ def __init__(
+ self,
+ patch_size: int = 2,
+ in_channels: int = 128,
+ out_channels: Optional[int] = 32,
+ num_layers: int = 48,
+ attention_head_dim: int = 64,
+ num_attention_heads: int = 24,
+ inner_dim: int = 1536, # pylint: disable=unused-argument
+ enc_hidden_dim: int = 2880,
+ axes_dims_rope: Tuple[int, int, int] = (8, 28, 28),
+ gate_mlp: bool = True,
+ rms_norm: bool = True,
+ multi_layer_encoder_feature: bool = True,
+ selected_layer_index: Tuple[int, ...] = (5, 11, 17, 23),
+ ) -> None:
+ super().__init__()
+ self.in_channels = in_channels
+ self.out_channels = out_channels or in_channels
+ self.inner_dim = num_attention_heads * attention_head_dim
+ self.multi_layer_encoder_feature = multi_layer_encoder_feature
+ self.selected_layer_index = list(selected_layer_index)
+
+ self.pos_embed = LensEmbedRope(theta=10000, axes_dim=list(axes_dims_rope), scale_rope=True)
+ self.time_text_embed = LensTimestepProjEmbeddings(embedding_dim=self.inner_dim)
+
+ if self.multi_layer_encoder_feature:
+ self.txt_norm = nn.ModuleList(
+ [RMSNorm(enc_hidden_dim, eps=1e-5) for _ in self.selected_layer_index]
+ )
+ self.txt_in = nn.Linear(enc_hidden_dim * len(self.selected_layer_index), self.inner_dim)
+ else:
+ self.txt_norm = RMSNorm(enc_hidden_dim, eps=1e-5)
+ self.txt_in = nn.Linear(enc_hidden_dim, self.inner_dim)
+
+ self.img_in = nn.Linear(in_channels, self.inner_dim)
+
+ self.transformer_blocks = nn.ModuleList(
+ [
+ LensTransformerBlock(
+ dim=self.inner_dim,
+ num_attention_heads=num_attention_heads,
+ attention_head_dim=attention_head_dim,
+ rms_norm=rms_norm,
+ gate_mlp=gate_mlp,
+ )
+ for _ in range(num_layers)
+ ]
+ )
+ self.norm_out = AdaLayerNormContinuous(
+ self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6
+ )
+ self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
+ encoder_hidden_states_mask: torch.Tensor,
+ timestep: torch.Tensor,
+ img_shapes: List[Tuple[int, int, int]],
+ attention_kwargs: Optional[Dict[str, Any]] = None, # pylint: disable=unused-argument
+ ) -> torch.Tensor:
+ """Forward pass.
+
+ Args:
+ hidden_states: [B, S_img, in_channels] image latents.
+ encoder_hidden_states: either a Tensor [B, S_txt, enc_dim]
+ (single-layer) or a list of such
+ tensors (multi-layer).
+ encoder_hidden_states_mask: bool [B, S_txt] (True = valid).
+ timestep: [B] in [0, 1].
+ img_shapes: list with a single (frame, h_lat, w_lat).
+ """
+ bsz, img_len, _ = hidden_states.shape
+ if self.multi_layer_encoder_feature:
+ if not isinstance(encoder_hidden_states, (list, tuple)):
+ raise ValueError(
+ "multi_layer_encoder_feature=True expects a list of "
+ "per-layer text tensors."
+ )
+ if len(encoder_hidden_states) != len(self.selected_layer_index):
+ raise ValueError(
+ f"Expected {len(self.selected_layer_index)} text feature "
+ f"layers, got {len(encoder_hidden_states)}."
+ )
+ text_seq_len = encoder_hidden_states[0].shape[1]
+ for i, feat in enumerate(encoder_hidden_states):
+ if feat.shape[0] != bsz:
+ raise ValueError(
+ f"Text feature layer {i} batch size {feat.shape[0]} "
+ f"does not match hidden_states batch size {bsz}."
+ )
+ if feat.shape[1] != text_seq_len:
+ raise ValueError(
+ f"Text feature layer {i} sequence length {feat.shape[1]} "
+ f"does not match layer 0 length {text_seq_len}."
+ )
+ else:
+ if not isinstance(encoder_hidden_states, torch.Tensor):
+ raise ValueError(
+ "multi_layer_encoder_feature=False expects a single text "
+ "feature tensor."
+ )
+ if encoder_hidden_states.shape[0] != bsz:
+ raise ValueError(
+ f"Text feature batch size {encoder_hidden_states.shape[0]} "
+ f"does not match hidden_states batch size {bsz}."
+ )
+ text_seq_len = encoder_hidden_states.shape[1]
+ if encoder_hidden_states_mask.shape != (bsz, text_seq_len):
+ raise ValueError(
+ "encoder_hidden_states_mask must have shape "
+ f"{(bsz, text_seq_len)}, got {tuple(encoder_hidden_states_mask.shape)}."
+ )
+ attention_mask = self._build_joint_attention_mask(
+ encoder_hidden_states_mask, img_len
+ )
+
+ hidden_states = self.img_in(hidden_states)
+ timestep = timestep.to(hidden_states.dtype)
+
+ if self.multi_layer_encoder_feature:
+ normed = [
+ self.txt_norm[i](encoder_hidden_states[i])
+ for i in range(len(self.selected_layer_index))
+ ]
+ encoder_hidden_states = torch.cat(normed, dim=-1)
+ else:
+ encoder_hidden_states = self.txt_norm(encoder_hidden_states)
+ encoder_hidden_states = self.txt_in(encoder_hidden_states)
+
+ temb = self.time_text_embed(timestep, hidden_states)
+
+ image_rotary_emb = self.pos_embed(
+ img_shapes, [text_seq_len], device=hidden_states.device
+ )
+
+ for block in self.transformer_blocks:
+ encoder_hidden_states, hidden_states = block(
+ hidden_states=hidden_states,
+ encoder_hidden_states=encoder_hidden_states,
+ temb=temb,
+ image_rotary_emb=image_rotary_emb,
+ attention_mask=attention_mask,
+ )
+
+ hidden_states = self.norm_out(hidden_states, temb)
+ return self.proj_out(hidden_states)
+
+ @staticmethod
+ def _build_joint_attention_mask(
+ text_mask: torch.Tensor, img_len: int
+ ) -> torch.Tensor:
+ """Additive joint mask of shape ``[B, 1, 1, img_len + S_txt]``.
+
+ Image tokens are always valid; text positions follow ``text_mask``.
+ Padded positions hold ``-inf`` so SDPA's softmax masks them out.
+ """
+ if text_mask.dtype != torch.bool:
+ text_mask = text_mask.bool()
+ bsz = text_mask.shape[0]
+ img_ones = torch.ones(
+ (bsz, img_len), dtype=torch.bool, device=text_mask.device
+ )
+ joint = torch.cat([img_ones, text_mask], dim=1)
+ additive = torch.zeros_like(joint, dtype=torch.float32)
+ additive.masked_fill_(~joint, float("-inf"))
+ return additive[:, None, None, :]
diff --git a/pipelines/model_lens.py b/pipelines/model_lens.py
new file mode 100644
index 000000000..8c422eb7d
--- /dev/null
+++ b/pipelines/model_lens.py
@@ -0,0 +1,38 @@
+from modules import shared, devices, sd_models, model_quant, sd_hijack_te, sd_hijack_vae
+from modules.logger import log
+from pipelines import generic
+
+
+def load_lens(checkpoint_info, diffusers_load_config=None):
+ if diffusers_load_config is None:
+ diffusers_load_config = {}
+
+ repo_id = sd_models.path_to_repo(checkpoint_info)
+ sd_models.hf_auth_check(checkpoint_info)
+
+ load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
+ log.debug(f'Load model: type=Lens repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
+
+ from pipelines import lens
+
+ transformer = generic.load_transformer(repo_id, cls_name=lens.LensTransformer2DModel, load_config=diffusers_load_config)
+ text_encoder = generic.load_text_encoder(repo_id, cls_name=lens.LensGptOssEncoder, load_config=diffusers_load_config, allow_quant=False) # te is prequantized using mxfp4
+
+ if not shared.opts.model_lens_enable_pe:
+ load_args['reasoner'] = None
+
+ pipe = lens.LensPipeline.from_pretrained(
+ repo_id,
+ transformer=transformer,
+ text_encoder=text_encoder,
+ cache_dir=shared.opts.diffusers_dir,
+ **load_args,
+ )
+ pipe.task_args = {
+ "output_type": "np",
+ }
+
+ sd_hijack_te.init_hijack(pipe)
+ sd_hijack_vae.init_hijack(pipe)
+ devices.torch_gc(force=True, reason="load")
+ return pipe