lens fixes

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2026-05-24 08:45:41 +02:00
parent 952f998d27
commit ded26d1fbc
5 changed files with 16 additions and 12 deletions
+1 -1
View File
@@ -4,7 +4,6 @@
- Inpaint: https://discord.com/channels/1101998836328697867/1130536562422186044/1506850651035144322
- Torch lazy-load
- pnpm-vs-npm
## Features
@@ -13,6 +12,7 @@
- Chat-based interface, @vladmandic
- Control tab verify overrides handling, @vladmandic
- Reimplement `llama` remover for Kanvas, @vladmandic
Object clear: https://huggingface.co/jixin0101/ObjectClear
- Detailer postprocessing, @CalamitousFelicitousness
- Cloud providers, @CalamitousFelicitousness
+3 -3
View File
@@ -131,7 +131,8 @@ class LensPipeline(DiffusionPipeline):
text_encoder: LensGptOssEncoder,
tokenizer: PreTrainedTokenizerBase,
transformer: LensTransformer2DModel,
reasoner: Optional[PromptReasoner] = True,
reasoner: Optional[PromptReasoner] = None,
use_reasoner=False,
) -> None:
super().__init__()
self.register_modules(
@@ -156,7 +157,7 @@ class LensPipeline(DiffusionPipeline):
self.transformer.config.selected_layer_index
)
if reasoner is not None:
if use_reasoner and reasoner is None:
self.reasoner = PromptReasoner(
text_encoder=self.text_encoder, tokenizer=self.tokenizer
)
@@ -296,7 +297,6 @@ class LensPipeline(DiffusionPipeline):
) -> List[str]:
if self.reasoner is None:
return list(prompts)
print('HERE REFINE')
return self.reasoner.refine(prompts, enable=enable_reasoner)
# ------------------------------------------------------------------
+9 -3
View File
@@ -38,11 +38,11 @@ class LensGptOssEncoder(GptOssForCausalLM):
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)
self._lens_selected_layers = layers # pylint: disable=attribute-defined-outside-init
self._lens_max_layer = max(layers) # pylint: disable=attribute-defined-outside-init
@torch.no_grad()
def forward( # type: ignore[override]
def forward( # type: ignore[override] # pylint: keyword-arg-before-vararg
self,
input_ids: Optional[torch.LongTensor] = None,
attention_mask: Optional[torch.Tensor] = None,
@@ -78,6 +78,12 @@ class LensGptOssEncoder(GptOssForCausalLM):
model = self.model
inputs_embeds = model.embed_tokens(input_ids)
target_dtype = inputs_embeds.dtype
if len(model.layers) > 0:
target_dtype = model.layers[0].self_attn.k_proj.weight.dtype
if inputs_embeds.dtype != target_dtype:
inputs_embeds = inputs_embeds.to(dtype=target_dtype)
position_ids = torch.arange(
inputs_embeds.shape[1], device=inputs_embeds.device
).unsqueeze(0).expand_as(input_ids)
+2 -4
View File
@@ -11,16 +11,14 @@ def load_lens(checkpoint_info, diffusers_load_config=None):
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}')
log.debug(f'Load model: type=Lens repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} reasoner={shared.opts.model_lens_enable_pe} 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
load_args['use_reasoner'] = shared.opts.model_lens_enable_pe
pipe = lens.LensPipeline.from_pretrained(
repo_id,
transformer=transformer,
+1 -1
View File
@@ -125,7 +125,6 @@ export async function getExif(el) {
// let html = `<b>Image</b> <a href="${el.src}" target="_blank">${el.src}</a> <b>Size</b> ${el.naturalWidth}x${el.naturalHeight}<br>`;
let html = '';
let params;
debug('getExif', exif);
if (exif.parameters) {
params = exif.parameters;
} else if (exif.userComment) {
@@ -133,6 +132,7 @@ export async function getExif(el) {
} else {
params = '';
}
debug('getExif', params);
if (params.length > 0) html += `<b>Prompt</b> ${params || ''}<br>`;
html = html.replace('Negative prompt:', '<br><b>Negative</b>');
html = html.replace('Steps:', '<br><b>Params</b> Steps:');