diff --git a/ttNpy/adv_encode.py b/ttNpy/adv_encode.py index c425c65..74fa172 100644 --- a/ttNpy/adv_encode.py +++ b/ttNpy/adv_encode.py @@ -6,10 +6,9 @@ from math import gcd from comfy import model_management from comfy.sdxl_clip import SDXLClipModel, SDXLRefinerClipModel, SDXLClipG, StableCascadeClipModel try: + from comfy.text_encoders.sd3_clip import SD3ClipModel, T5XXLModel +except ImportError: from comfy.sd3_clip import SD3ClipModel, T5XXLModel -except: - SD3ClipModel, T5XXLModel = None, None - pass try: from comfy.text_encoders.flux import FluxClipModel @@ -27,8 +26,8 @@ def _grouper(n, iterable): def _norm_mag(w, n): d = w - 1 - return 1 + np.sign(d) * np.sqrt(np.abs(d)**2 / n) - #return np.sign(w) * np.sqrt(np.abs(w)**2 / n) + return 1 + np.sign(d) * np.sqrt(np.abs(d) ** 2 / n) + # return np.sign(w) * np.sqrt(np.abs(w)**2 / n) def divide_length(word_ids, weights): sums = dict(zip(*np.unique(word_ids, return_counts=True))) @@ -38,28 +37,28 @@ def divide_length(word_ids, weights): return weights def shift_mean_weight(word_ids, weights): - delta = 1 - np.mean([w for x, y in zip(weights, word_ids) for w, id in zip(x,y) if id != 0]) - weights = [[w if id == 0 else w+delta + delta = 1 - np.mean([w for x, y in zip(weights, word_ids) for w, id in zip(x, y) if id != 0]) + weights = [[w if id == 0 else w + delta for w, id in zip(x, y)] for x, y in zip(weights, word_ids)] return weights def scale_to_norm(weights, word_ids, w_max): top = np.max(weights) w_max = min(top, w_max) - weights = [[w_max if id == 0 else (w/top) * w_max + weights = [[w_max if id == 0 else (w / top) * w_max for w, id in zip(x, y)] for x, y in zip(weights, word_ids)] return weights def from_zero(weights, base_emb): weight_tensor = torch.tensor(weights, dtype=base_emb.dtype, device=base_emb.device) - weight_tensor = weight_tensor.reshape(1,-1,1).expand(base_emb.shape) + weight_tensor = weight_tensor.reshape(1, -1, 1).expand(base_emb.shape) return base_emb * weight_tensor def mask_word_id(tokens, word_ids, target_id, mask_token): - new_tokens = [[mask_token if wid == target_id else t - for t, wid in zip(x,y)] for x,y in zip(tokens, word_ids)] - mask = np.array(word_ids) == target_id - return (new_tokens, mask) + new_tokens = [[mask_token if wid == target_id else t + for t, wid in zip(x, y)] for x, y in zip(tokens, word_ids)] + mask = np.array(word_ids) == target_id + return (new_tokens, mask) def batched_clip_encode(tokens, length, encode_func, num_chunks): embs = [] @@ -75,49 +74,49 @@ def batched_clip_encode(tokens, length, encode_func, num_chunks): return embs def from_masked(tokens, weights, word_ids, base_emb, length, encode_func, m_token=266): - pooled_base = base_emb[0,length-1:length,:] + pooled_base = base_emb[0, length - 1:length, :] wids, inds = np.unique(np.array(word_ids).reshape(-1), return_index=True) - weight_dict = dict((id,w) - for id,w in zip(wids ,np.array(weights).reshape(-1)[inds]) - if w != 1.0) + weight_dict = dict((id, w) + for id, w in zip(wids, np.array(weights).reshape(-1)[inds]) + if w != 1.0) if len(weight_dict) == 0: - return torch.zeros_like(base_emb), base_emb[0,length-1:length,:] + return torch.zeros_like(base_emb), base_emb[0, length - 1:length, :] weight_tensor = torch.tensor(weights, dtype=base_emb.dtype, device=base_emb.device) - weight_tensor = weight_tensor.reshape(1,-1,1).expand(base_emb.shape) + weight_tensor = weight_tensor.reshape(1, -1, 1).expand(base_emb.shape) - #m_token = (clip.tokenizer.end_token, 1.0) if clip.tokenizer.pad_with_end else (0,1.0) - #TODO: find most suitable masking token here + # m_token = (clip.tokenizer.end_token, 1.0) if clip.tokenizer.pad_with_end else (0,1.0) + # TODO: find most suitable masking token here m_token = (m_token, 1.0) ws = [] masked_tokens = [] masks = [] - #create prompts + # create prompts for id, w in weight_dict.items(): masked, m = mask_word_id(tokens, word_ids, id, m_token) masked_tokens.extend(masked) - + m = torch.tensor(m, dtype=base_emb.dtype, device=base_emb.device) - m = m.reshape(1,-1,1).expand(base_emb.shape) + m = m.reshape(1, -1, 1).expand(base_emb.shape) masks.append(m) ws.append(w) - - #batch process prompts + + # batch process prompts embs = batched_clip_encode(masked_tokens, length, encode_func, len(tokens)) masks = torch.cat(masks) - + embs = (base_emb.expand(embs.shape) - embs) - pooled = embs[0,length-1:length,:] + pooled = embs[0, length - 1:length, :] embs *= masks embs = embs.sum(axis=0, keepdim=True) pooled_start = pooled_base.expand(len(ws), -1) - ws = torch.tensor(ws).reshape(-1,1).expand(pooled_start.shape) + ws = torch.tensor(ws).reshape(-1, 1).expand(pooled_start.shape) pooled = (pooled - pooled_start) * (ws - 1) pooled = pooled.mean(axis=0, keepdim=True) @@ -126,17 +125,17 @@ def from_masked(tokens, weights, word_ids, base_emb, length, encode_func, m_toke def mask_inds(tokens, inds, mask_token): clip_len = len(tokens[0]) inds_set = set(inds) - new_tokens = [[mask_token if i*clip_len + j in inds_set else t + new_tokens = [[mask_token if i * clip_len + j in inds_set else t for j, t in enumerate(x)] for i, x in enumerate(tokens)] return new_tokens def down_weight(tokens, weights, word_ids, base_emb, length, encode_func, m_token=266): - w, w_inv = np.unique(weights,return_inverse=True) + w, w_inv = np.unique(weights, return_inverse=True) if np.sum(w < 1) == 0: - return base_emb, tokens, base_emb[0,length-1:length,:] - #m_token = (clip.tokenizer.end_token, 1.0) if clip.tokenizer.pad_with_end else (0,1.0) - #using the comma token as a masking token seems to work better than aos tokens for SD 1.x + return base_emb, tokens, base_emb[0, length - 1:length, :] + # m_token = (clip.tokenizer.end_token, 1.0) if clip.tokenizer.pad_with_end else (0,1.0) + # using the comma token as a masking token seems to work better than aos tokens for SD 1.x m_token = (m_token, 1.0) masked_tokens = [] @@ -150,12 +149,12 @@ def down_weight(tokens, weights, word_ids, base_emb, length, encode_func, m_toke embs = batched_clip_encode(masked_tokens, length, encode_func, len(tokens)) embs = torch.cat([base_emb, embs]) - w = w[w<=1.0] + w = w[w <= 1.0] w_mix = np.diff([0] + w.tolist()) - w_mix = torch.tensor(w_mix, dtype=embs.dtype, device=embs.device).reshape((-1,1,1)) + w_mix = torch.tensor(w_mix, dtype=embs.dtype, device=embs.device).reshape((-1, 1, 1)) weighted_emb = (w_mix * embs).sum(axis=0, keepdim=True) - return weighted_emb, masked_current, weighted_emb[0,length-1:length,:] + return weighted_emb, masked_current, weighted_emb[0, length - 1:length, :] def scale_emb_to_mag(base_emb, weighted_emb): norm_base = torch.linalg.norm(base_emb) @@ -172,48 +171,49 @@ def A1111_renorm(base_emb, weighted_emb): embeddings_final = (base_emb.mean() / weighted_emb.mean()) * weighted_emb return embeddings_final -def advanced_encode_from_tokens(tokenized, token_normalization, weight_interpretation, encode_func, m_token=266, length=77, w_max=1.0, return_pooled=False, apply_to_pooled=False): - tokens = [[t for t,_,_ in x] for x in tokenized] - weights = [[w for _,w,_ in x] for x in tokenized] - word_ids = [[wid for _,_,wid in x] for x in tokenized] +def advanced_encode_from_tokens(tokenized, token_normalization, weight_interpretation, encode_func, m_token=266, + length=77, w_max=1.0, return_pooled=False, apply_to_pooled=False): + tokens = [[t for t, _, _ in x] for x in tokenized] + weights = [[w for _, w, _ in x] for x in tokenized] + word_ids = [[wid for _, _, wid in x] for x in tokenized] - #weight normalization - #==================== + # weight normalization + # ==================== - #distribute down/up weights over word lengths + # distribute down/up weights over word lengths if token_normalization.startswith("length"): weights = divide_length(word_ids, weights) - - #make mean of word tokens 1 - if token_normalization.endswith("mean"): - weights = shift_mean_weight(word_ids, weights) - #weight interpretation - #===================== + # make mean of word tokens 1 + if token_normalization.endswith("mean"): + weights = shift_mean_weight(word_ids, weights) + + # weight interpretation + # ===================== pooled = None if weight_interpretation == "comfy": - weighted_tokens = [[(t,w) for t, w in zip(x, y)] for x, y in zip(tokens, weights)] + weighted_tokens = [[(t, w) for t, w in zip(x, y)] for x, y in zip(tokens, weights)] weighted_emb, pooled_base = encode_func(weighted_tokens) pooled = pooled_base else: - unweighted_tokens = [[(t,1.0) for t, _,_ in x] for x in tokenized] + unweighted_tokens = [[(t, 1.0) for t, _, _ in x] for x in tokenized] base_emb, pooled_base = encode_func(unweighted_tokens) - + if weight_interpretation == "A1111": weighted_emb = from_zero(weights, base_emb) weighted_emb = A1111_renorm(base_emb, weighted_emb) pooled = pooled_base - + if weight_interpretation == "compel": - pos_tokens = [[(t,w) if w >= 1.0 else (t,1.0) for t, w in zip(x, y)] for x, y in zip(tokens, weights)] + pos_tokens = [[(t, w) if w >= 1.0 else (t, 1.0) for t, w in zip(x, y)] for x, y in zip(tokens, weights)] weighted_emb, _ = encode_func(pos_tokens) weighted_emb, _, pooled = down_weight(pos_tokens, weights, word_ids, weighted_emb, length, encode_func) - + if weight_interpretation == "comfy++": weighted_emb, tokens_down, _ = down_weight(unweighted_tokens, weights, word_ids, base_emb, length, encode_func) weights = [[w if w > 1.0 else 1.0 for w in x] for x in weights] - #unweighted_tokens = [[(t,1.0) for t, _,_ in x] for x in tokens_down] + # unweighted_tokens = [[(t,1.0) for t, _, _ in x] for x in tokens_down] embs, pooled = from_masked(unweighted_tokens, weights, word_ids, base_emb, length, encode_func) weighted_emb += embs @@ -232,8 +232,7 @@ def encode_token_weights_g(model, token_weight_pairs): return model.clip_g.encode_token_weights(token_weight_pairs) def encode_token_weights_l(model, token_weight_pairs): - l_out, l_pooled = model.clip_l.encode_token_weights(token_weight_pairs) - return l_out, l_pooled + return model.clip_l.encode_token_weights(token_weight_pairs) def encode_token_weights_t5(model, token_weight_pairs): return model.t5xxl.encode_token_weights(token_weight_pairs) @@ -241,7 +240,7 @@ def encode_token_weights_t5(model, token_weight_pairs): def encode_token_weights(model, token_weight_pairs, encode_func): if model.layer_idx is not None: model.cond_stage_model.set_clip_options({"layer": model.layer_idx}) - + model_management.load_model_gpu(model.patcher) return encode_func(model.cond_stage_model, token_weight_pairs) @@ -263,7 +262,7 @@ def prepareSD3(out, pooled, clip_balance): def advanced_encode(clip, text, token_normalization, weight_interpretation, w_max=1.0, clip_balance=.5, apply_to_pooled=True): tokenized = clip.tokenize(text, return_word_ids=True) - + if SD3ClipModel and isinstance(clip.cond_stage_model, SD3ClipModel): lg_out = None pooled = None @@ -301,10 +300,10 @@ def advanced_encode(clip, text, token_normalization, weight_interpretation, w_ma # t5xxl if 't5xxl' in tokenized and clip.cond_stage_model.t5xxl is not None: t5_out, t5_pooled = advanced_encode_from_tokens(tokenized['t5xxl'], - token_normalization, - weight_interpretation, - lambda x: encode_token_weights(clip, x, encode_token_weights_t5), - w_max=w_max, return_pooled=True) + token_normalization, + weight_interpretation, + lambda x: encode_token_weights(clip, x, encode_token_weights_t5), + w_max=w_max, return_pooled=True) if lg_out is not None: out = torch.cat([lg_out, t5_out], dim=-2) else: @@ -343,18 +342,18 @@ def advanced_encode(clip, text, token_normalization, weight_interpretation, w_ma embs_g = None pooled = None if 'l' in tokenized and isinstance(clip.cond_stage_model, SDXLClipModel): - embs_l, _ = advanced_encode_from_tokens(tokenized['l'], - token_normalization, - weight_interpretation, + embs_l, _ = advanced_encode_from_tokens(tokenized['l'], + token_normalization, + weight_interpretation, lambda x: encode_token_weights(clip, x, encode_token_weights_l), - w_max=w_max, + w_max=w_max, return_pooled=False) if 'g' in tokenized: - embs_g, pooled = advanced_encode_from_tokens(tokenized['g'], - token_normalization, + embs_g, pooled = advanced_encode_from_tokens(tokenized['g'], + token_normalization, weight_interpretation, lambda x: encode_token_weights(clip, x, encode_token_weights_g), - w_max=w_max, + w_max=w_max, return_pooled=True, apply_to_pooled=apply_to_pooled) return prepareXL(embs_l, embs_g, pooled, clip_balance) diff --git a/ttNpy/tinyterraNodes.py b/ttNpy/tinyterraNodes.py index 4153812..f5a1bde 100644 --- a/ttNpy/tinyterraNodes.py +++ b/ttNpy/tinyterraNodes.py @@ -158,8 +158,8 @@ class ttNloader: clip = loaded_ckpt[1].clone() if loaded_ckpt[1] is not None else None if clip_skip != 0 and clip is not None: - if sampler.get_model_type(loaded_ckpt[0]) == 'FLUX': - raise Exception('Flux does not support clip_skip. Set clip_skip to 0.') + if sampler.get_model_type() in ['FLUX', 'FLOW']: + raise Exception('FLOW and FLUX do not support clip_skip. Set clip_skip to 0.') clip.clip_layer(clip_skip) # model, clip, vae @@ -327,8 +327,8 @@ class ttNloader: clip = clip_override.clone() if clip_skip != 0: - if sampler.get_model_type() == 'FLUX': - raise Exception('Flux does not support clip_skip. Set clip_skip to 0.') + if sampler.get_model_type() in ['FLUX', 'FLOW']: + raise Exception('FLOW and FLUX do not support clip_skip. Set clip_skip to 0.') clip.clip_layer(clip_skip) del clip_override @@ -1183,7 +1183,7 @@ class ttN_pipeLoader_v2: model, clip, vae = loader.load_main3(ckpt_name, config_name, vae_name, loras, clip_skip, model_override, clip_override, optional_lora_stack, my_unique_id) # Create Empty Latent - sd3 = True if sampler.get_model_type(model) in ['FLUX', 'SD3'] else False + sd3 = True if sampler.get_model_type(model) in ['FLUX', 'FLOW'] else False latent = sampler.emptyLatent(empty_latent_aspect, batch_size, empty_latent_width, empty_latent_height, sd3) samples = {"samples":latent} @@ -1580,7 +1580,7 @@ class ttN_pipeLoaderSDXL_v2: model, clip, vae = loader.load_main3(ckpt_name, config_name, vae_name, loras, clip_skip, model_override, clip_override, optional_lora_stack, my_unique_id) # Create Empty Latent - sd3 = True if sampler.get_model_type(model) in ['FLUX', 'SD3'] else False + sd3 = True if sampler.get_model_type(model) in ['FLUX', 'FLOW'] else False latent = sampler.emptyLatent(empty_latent_aspect, batch_size, empty_latent_width, empty_latent_height, sd3) samples = {"samples":latent} @@ -2268,7 +2268,7 @@ class ttN_tinyLoader: model, clip, vae = loader.load_checkpoint(ckpt_name, config_name, clip_skip) # Create Empty Latent - sd3 = True if sampler.get_model_type(model) in ['FLUX', 'SD3'] else False + sd3 = True if sampler.get_model_type(model) in ['FLUX', 'FLOW'] else False latent = sampler.emptyLatent(empty_latent_aspect, batch_size, empty_latent_width, empty_latent_height, sd3) samples = {"samples": latent}