ip adapter masking

This commit is contained in:
Vladimir Mandic
2024-04-18 12:43:51 -04:00
parent 77909a103a
commit f2610c3936
12 changed files with 84 additions and 31 deletions
+10 -2
View File
@@ -2,15 +2,22 @@
## Pending
- PixArt-Σ requires `diffusers-0.28.0.dev0`
### requires `diffusers-0.28.0.dev0`
## Update for 2024-04-17
- PixArt-Σ
- IP adapter masking
## Update for 2024-04-18
- **Features**:
- **Gallery**: list, preview, search through all your images and videos!
implemented as infinite-scroll with client-side-caching and lazy-loading while being fully async and non-blocking
search or sort by path, name, size, width, height, mtime or any image metadata item, also with extended syntax like *width > 1000*
*settings*: optional additional user-defined folders, thumbnails in fixed or variable aspect-ratio
- **IP Adapter Masking**:
powerful method of using masking with ip-adapters
when combined with multiple ip-adapters, it allows for different inputs guidance for each segment of the input image
*hint*: to create masks, you can use manually created masks or control->mask module with auto-segment to create masks and later upload them
- **OneDiff**: new optimization/compile engine, thanks @aifartist
as with all other compile engines, enable via *settings -> compute settings -> compile*
- **UI**:
@@ -47,6 +54,7 @@
- Faster server startup
- Add **MIGraphX** torch optimization engine, thanks @Disty0
- Styles apply wildcards to params
- Extra networks persistent sort order in settings
- Add option to make batch generations use fully random seed vs sequential
- Make metadata in full screen viewer optional
- Add VAE civitai scan metadata/preview
+7 -4
View File
@@ -14,13 +14,11 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma
- stable diffusion 3.0
- powerpaint: <https://github.com/zhuang2002/PowerPaint>
- ella: <https://github.com/TencentQQGYLab/ELLA>
### Pipelines
- instant style: <https://github.com/huggingface/diffusers/pull/7586> <https://github.com/InstantStyle/InstantStyle>
- instant style: <https://github.com/huggingface/diffusers/pull/7668> <https://github.com/InstantStyle/InstantStyle>
- ipadapter masking: <https://github.com/huggingface/diffusers/pull/6847>
- x-adapter: <https://github.com/showlab/X-Adapter>
### Features
@@ -30,8 +28,13 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma
- include reference styles
- quick apply style
- lora: sc lora, dora, etc
- samplers: smea <https://github.com/Koishi-Star/Euler-Smea-Dyn-Sampler>, restart <https://github.com/Newbeeer/diffusion_restart_sampling>
### Missing
- control api scripts compatibility
### Defer
- ella: <https://github.com/TencentQQGYLab/ELLA>
- x-adapter: <https://github.com/showlab/X-Adapter>
- samplers: smea <https://github.com/Koishi-Star/Euler-Smea-Dyn-Sampler>, restart <https://github.com/Newbeeer/diffusion_restart_sampling>
Binary file not shown.

After

Width:  |  Height:  |  Size: 196 KiB

+23 -18
View File
@@ -1,5 +1,5 @@
const activePromptTextarea = {};
let sortVal = 0;
let sortVal = -1;
// helpers
@@ -224,33 +224,35 @@ function tryToRemoveExtraNetworkFromPrompt(textarea, text) {
return false;
}
function sortExtraNetworks() {
const sortDesc = ['Name [A-Z]', 'Name [Z-A]', 'Date [Newest]', 'Date [Oldest]', 'Size [Largest]', 'Size [Smallest]'];
function sortExtraNetworks(fixed = 'no') {
const sortDesc = ['Default', 'Name [A-Z]', 'Name [Z-A]', 'Date [Newest]', 'Date [Oldest]', 'Size [Largest]', 'Size [Smallest]'];
const pagename = getENActivePage();
if (!pagename) return 'sort error: unknown page';
const allPages = Array.from(gradioApp().querySelectorAll('.extra-network-cards'));
const pages = allPages.filter((el) => el.id.toLowerCase().includes(pagename.toLowerCase()));
let num = 0;
if (sortVal === -1) sortVal = sortDesc.indexOf(opts.extra_networks_sort);
if (fixed !== 'fixed') sortVal = (sortVal + 1) % sortDesc.length;
for (const pg of pages) {
const cards = Array.from(pg.querySelectorAll('.card') || []);
num = cards.length;
if (num === 0) return 'sort: no cards';
cards.sort((a, b) => { // eslint-disable-line no-loop-func
switch (sortVal) {
case 0: return a.dataset.name ? a.dataset.name.localeCompare(b.dataset.name) : 0;
case 1: return b.dataset.name ? b.dataset.name.localeCompare(a.dataset.name) : 0;
case 2: return a.dataset.mtime && !isNaN(a.dataset.mtime) ? parseFloat(b.dataset.mtime) - parseFloat(a.dataset.mtime) : 0;
case 3: return b.dataset.mtime && !isNaN(b.dataset.mtime) ? parseFloat(a.dataset.mtime) - parseFloat(b.dataset.mtime) : 0;
case 4: return a.dataset.size && !isNaN(a.dataset.size) ? parseFloat(b.dataset.size) - parseFloat(a.dataset.size) : 0;
case 5: return b.dataset.size && !isNaN(b.dataset.size) ? parseFloat(a.dataset.size) - parseFloat(b.dataset.size) : 0;
case 0: return 0;
case 1: return a.dataset.name ? a.dataset.name.localeCompare(b.dataset.name) : 0;
case 2: return b.dataset.name ? b.dataset.name.localeCompare(a.dataset.name) : 0;
case 3: return a.dataset.mtime && !isNaN(a.dataset.mtime) ? parseFloat(b.dataset.mtime) - parseFloat(a.dataset.mtime) : 0;
case 4: return b.dataset.mtime && !isNaN(b.dataset.mtime) ? parseFloat(a.dataset.mtime) - parseFloat(b.dataset.mtime) : 0;
case 5: return a.dataset.size && !isNaN(a.dataset.size) ? parseFloat(b.dataset.size) - parseFloat(a.dataset.size) : 0;
case 6: return b.dataset.size && !isNaN(b.dataset.size) ? parseFloat(a.dataset.size) - parseFloat(b.dataset.size) : 0;
}
return 0;
});
for (const card of cards) pg.appendChild(card);
}
const desc = sortDesc[sortVal];
sortVal = (sortVal + 1) % sortDesc.length;
log('sortExtraNetworks', pagename, num, desc);
log('sortExtraNetworks', { name: pagename, val: sortVal, order: desc, fixed: fixed === 'fixed', items: num });
return `sort page ${pagename} cards ${num} by ${desc}`;
}
@@ -271,12 +273,8 @@ function extraNetworksSearchButton(event) {
const tabname = getENActiveTab();
const searchTextarea = gradioApp().querySelector(`#${tabname}_extra_search textarea`);
const button = event.target;
if (button.classList.contains('search-all')) {
searchTextarea.value = '';
} else {
searchTextarea.value = `${button.textContent.trim()}/`;
}
if (button.classList.contains('search-all')) searchTextarea.value = '';
else searchTextarea.value = `${button.textContent.trim()}/`;
updateInput(searchTextarea);
}
@@ -307,6 +305,11 @@ function quickSaveStyle() {
const tabname = getENActiveTab();
const btnSave = gradioApp().getElementById(`${tabname}_extra_quicksave`);
if (btnSave) btnSave.click();
const btnRefresh = gradioApp().getElementById(`${tabname}_extra_refresh`);
if (btnRefresh) {
setTimeout(() => btnRefresh.click(), 100);
setTimeout(() => sortExtraNetworks('fixed'), 500);
}
}
let enDirty = false;
@@ -368,6 +371,7 @@ function setupExtraNetworksForTab(tabname) {
if (btnView) buttons.appendChild(btnView);
if (btnClose) buttons.appendChild(btnClose);
btnModel.onclick = () => btnModel.classList.toggle('toolbutton-selected');
btnRefresh.onclick = () => sortExtraNetworks('fixed');
tabs.appendChild(buttons);
// details
@@ -426,6 +430,7 @@ function setupExtraNetworksForTab(tabname) {
}
if (entries[0].intersectionRatio > 0) {
refreshENpage();
sortExtraNetworks('fixed');
if (window.opts.extra_networks_card_cover === 'cover') {
en.style.transition = '';
en.style.zIndex = 100;
@@ -459,7 +464,7 @@ function setupExtraNetworksForTab(tabname) {
gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = 'unset';
}
});
intersectionObserver.observe(en); // monitor visibility of
intersectionObserver.observe(en); // monitor visibility
}
async function setupExtraNetworks() {
+1 -1
View File
@@ -213,7 +213,7 @@ table.settings-value-table td { padding: 0.4em; border: 1px solid #ccc; max-widt
.extra-network-cards { display: flex; flex-wrap: wrap; overflow-y: auto; overflow-x: hidden; align-content: flex-start; width: -moz-available; width: -webkit-fill-available; }
.extra-network-cards .card { height: fit-content; margin: 0 0 0.5em 0.5em; position: relative; scroll-snap-align: start; scroll-margin-top: 0; }
.extra-network-cards .card .overlay { z-index: 10; width: 100%; background: none; }
.extra-network-cards .card .overlay .name { font-size: var(--text-lg); font-weight: bold; text-shadow: 1px 1px black; color: white; overflow-wrap: break-word; position: absolute; bottom: 0; padding: 0.2em; z-index: 10; }
.extra-network-cards .card .overlay .name { font-size: var(--text-lg); font-weight: bold; text-shadow: 1px 1px black; color: white; overflow-wrap: anywhere; position: absolute; bottom: 0; padding: 0.2em; z-index: 10; }
.extra-network-cards .card .preview { box-shadow: var(--button-shadow); min-height: 30px; }
.extra-network-cards .card:hover .overlay { background: rgba(0, 0, 0, 0.70); }
.extra-network-cards .card:hover .preview { box-shadow: none; filter: grayscale(100%); }
+2
View File
@@ -67,6 +67,7 @@ class APIGenerate():
p.ip_adapter_starts = []
p.ip_adapter_ends = []
p.ip_adapter_images = []
p.ip_adapter_masks = []
for ipadapter in request.ip_adapter:
if not ipadapter.images or len(ipadapter.images) == 0:
continue
@@ -75,6 +76,7 @@ class APIGenerate():
p.ip_adapter_starts.append(ipadapter.start)
p.ip_adapter_ends.append(ipadapter.end)
p.ip_adapter_images.append([helpers.decode_base64_to_image(x) for x in ipadapter.images])
p.ip_adapter_masks.append([helpers.decode_base64_to_image(x) for x in ipadapter.masks])
del request.ip_adapter
def post_text2img(self, txt2imgreq: models.ReqTxt2Img):
+1
View File
@@ -151,6 +151,7 @@ class ItemEmbedding(BaseModel):
class ItemIPAdapter(BaseModel):
adapter: str = Field(title="Adapter", default="Base", description="")
images: List[str] = Field(title="Image", default=[], description="")
masks: Optional[List[str]] = Field(title="Mask", default=[], description="")
scale: float = Field(title="Scale", default=0.5, gt=0, le=1, description="")
start: float = Field(title="Start", default=0.0, gt=0, le=1, description="")
end: float = Field(title="End", default=1.0, gt=0, le=1, description="")
+17 -1
View File
@@ -108,8 +108,22 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_names=[], adapt
if hasattr(p, 'ip_adapter_images'):
adapter_images = p.ip_adapter_images
adapter_images = get_images(adapter_images)
if hasattr(p, 'ip_adapter_masks'):
adapter_masks = p.ip_adapter_masks
adapter_masks = get_images(adapter_masks)
if len(adapter_masks) > 0:
from diffusers.image_processor import IPAdapterMaskProcessor
mask_processor = IPAdapterMaskProcessor()
for i in range(len(adapter_masks)):
adapter_masks[i] = mask_processor.preprocess(adapter_masks[i], height=p.height, width=p.width)
adapter_masks = mask_processor.preprocess(adapter_masks, height=p.height, width=p.width)
if len(adapters) < len(adapter_images):
adapter_images = adapter_images[:len(adapters)]
if len(adapters) < len(adapter_masks):
adapter_masks = adapter_masks[:len(adapters)]
if len(adapter_masks) > 0 and len(adapter_masks) != len(adapter_images):
shared.log.error('IP adapter: image and mask count mismatch')
return False
adapter_scales = get_scales(adapter_scales, adapter_images)
p.ip_adapter_scales = adapter_scales.copy()
adapter_starts = get_scales(adapter_starts, adapter_images)
@@ -179,10 +193,12 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_names=[], adapt
adapter_scales[i] = 0.00
pipe.set_ip_adapter_scale(adapter_scales)
p.task_args['ip_adapter_image'] = adapter_images
if len(adapter_masks) > 0:
p.cross_attention_kwargs = { 'ip_adapter_masks': adapter_masks }
t1 = time.time()
ip_str = [f'{os.path.splitext(adapter)[0]}:{scale}:{start}:{end}' for adapter, scale, start, end in zip(adapter_names, adapter_scales, adapter_starts, adapter_ends)]
p.extra_generation_params["IP Adapter"] = ';'.join(ip_str)
shared.log.info(f'IP adapter: {ip_str} image={adapter_images} time={t1-t0:.2f}')
shared.log.info(f'IP adapter: {ip_str} image={adapter_images} mask={adapter_masks is not None} time={t1-t0:.2f}')
except Exception as e:
shared.log.error(f'IP adapter failed to load: repo={base_repo} folder={ip_subfolder} weights={adapters} names={adapter_names} {e}')
return True
+7
View File
@@ -282,6 +282,12 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
args[k] = v
else:
debug(f'Diffusers unknown task args: {k}={v}')
cross_attention_args = getattr(p, 'cross_attention_kwargs', {})
debug(f'Diffusers cross-attention args: {cross_attention_args}')
for k, v in cross_attention_args.items():
if args.get('cross_attention_kwargs', None) is None:
args['cross_attention_kwargs'] = {}
args['cross_attention_kwargs'][k] = v
# handle implicit controlnet
if 'control_image' in possible and 'control_image' not in args and 'image' in args:
@@ -292,6 +298,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
# debug info
clean = args.copy()
clean.pop('cross_attention_kwargs', None)
clean.pop('callback', None)
clean.pop('callback_steps', None)
clean.pop('callback_on_step_end', None)
+1
View File
@@ -759,6 +759,7 @@ options_templates.update(options_section(('interrogate', "Interrogate"), {
options_templates.update(options_section(('extra_networks', "Extra Networks"), {
"extra_networks_sep1": OptionInfo("<h2>Extra networks UI</h2>", "", gr.HTML),
"extra_networks": OptionInfo(["All"], "Extra networks", gr.Dropdown, lambda: {"multiselect":True, "choices": ['All'] + [en.title for en in extra_networks]}),
"extra_networks_sort": OptionInfo("Default", "Sort order", gr.Dropdown, {"choices": ['Default', 'Name [A-Z]', 'Name [Z-A]', 'Date [Newest]', 'Date [Oldest]', 'Size [Largest]', 'Size [Smallest]']}),
"extra_networks_view": OptionInfo("gallery", "UI view", gr.Radio, {"choices": ["gallery", "list"]}),
"extra_networks_card_cover": OptionInfo("sidebar", "UI position", gr.Radio, {"choices": ["cover", "inline", "sidebar"]}),
"extra_networks_height": OptionInfo(53, "UI height (%)", gr.Slider, {"minimum": 10, "maximum": 100, "step": 1}),
+1
View File
@@ -831,6 +831,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False):
for page in get_pages():
if title is None or title == '' or title == page.title or len(page.html) == 0:
shared.opts.extra_networks_view = page.view
# shared.opts.save(shared.config_filename)
page.view = 'gallery' if page.view == 'list' else 'list'
page.card = card_full if page.view == 'gallery' else card_list
page.html = ''
+14 -5
View File
@@ -33,7 +33,7 @@ class Script(scripts.Script):
init_images.append(image)
except Exception as e:
shared.log.warning(f'IP adapter failed to load image: {e}')
return init_images
return gr.update(value=init_images, visible=len(init_images) > 0)
def display_units(self, num_units):
num_units = num_units or 1
@@ -47,7 +47,9 @@ class Script(scripts.Script):
starts = []
ends = []
files = []
galleries = []
masks = []
image_galleries = []
mask_galleries = []
with gr.Row():
num_adapters = gr.Slider(label="Active IP adapters", minimum=1, maximum=MAX_ADAPTERS, step=1, value=1, scale=1)
for i in range(MAX_ADAPTERS):
@@ -61,11 +63,16 @@ class Script(scripts.Script):
with gr.Row():
files.append(gr.File(label='Input images', file_count='multiple', file_types=['image'], type='file', interactive=True, height=100))
with gr.Row():
galleries.append(gr.Gallery(show_label=False, value=[]))
files[i].change(fn=self.load_images, inputs=[files[i]], outputs=[galleries[i]])
image_galleries.append(gr.Gallery(show_label=False, value=[], visible=False, container=False, rows=1))
with gr.Row():
masks.append(gr.File(label='Input masks', file_count='multiple', file_types=['image'], type='file', interactive=True, height=100))
with gr.Row():
mask_galleries.append(gr.Gallery(show_label=False, value=[], visible=False))
files[i].change(fn=self.load_images, inputs=[files[i]], outputs=[image_galleries[i]])
masks[i].change(fn=self.load_images, inputs=[masks[i]], outputs=[mask_galleries[i]])
units.append(unit)
num_adapters.change(fn=self.display_units, inputs=[num_adapters], outputs=units)
return [num_adapters] + adapters + scales + files + starts + ends
return [num_adapters] + adapters + scales + files + starts + ends + masks
def process(self, p: processing.StableDiffusionProcessing, *args): # pylint: disable=arguments-differ
if shared.backend != shared.Backend.DIFFUSERS:
@@ -84,4 +91,6 @@ class Script(scripts.Script):
p.ip_adapter_starts = args[MAX_ADAPTERS*3:MAX_ADAPTERS*4][:units]
if getattr(p, 'ip_adapter_ends', [1.0]) == [1.0]:
p.ip_adapter_ends = args[MAX_ADAPTERS*4:MAX_ADAPTERS*5][:units]
if getattr(p, 'ip_adapter_masks', []) == []:
p.ip_adapter_masks = args[MAX_ADAPTERS*5:MAX_ADAPTERS*6][:units]
# ipadapter.apply(shared.sd_model, p, p.ip_adapter_names, p.ip_adapter_scales, p.ip_adapter_starts, p.ip_adapter_ends, p.ip_adapter_images) # called directly from processing.process_images_inner