Ported ContextRef adain support to work with uuids, reorganized the ContextRef-related code in BankStyle classes

This commit is contained in:
Jedrzej Kosinski
2024-10-11 09:43:06 -05:00
parent b0044cc2b7
commit 7d25e4bafb
+94 -69
View File
@@ -426,23 +426,11 @@ class BankStylesBasicTransformerBlock:
self.c_style_cfgs: dict[UUID, list[float]] = {} self.c_style_cfgs: dict[UUID, list[float]] = {}
self.c_cn_idx: dict[UUID, list[int]] = {} self.c_cn_idx: dict[UUID, list[int]] = {}
# self.c_bank: list[list] = []
# self.c_style_cfgs: list[list] = []
# self.c_cn_idx: list[list[int]] = []
def set_c_bank_for_uuids(self, x: Tensor, uuids: list[UUID]): def set_c_bank_for_uuids(self, x: Tensor, uuids: list[UUID]):
per_uuid = len(x) // len(uuids) per_uuid = len(x) // len(uuids)
for uuid, i in zip(uuids, list(range(0, len(x), per_uuid))): for uuid, i in zip(uuids, list(range(0, len(x), per_uuid))):
self.c_bank.setdefault(uuid, []).append(x[i:i+per_uuid]) self.c_bank.setdefault(uuid, []).append(x[i:i+per_uuid])
def set_c_style_cfgs_for_uuids(self, style_cfg: float, uuids: list[UUID]):
for uuid in uuids:
self.c_style_cfgs.setdefault(uuid, []).append(style_cfg)
def set_c_cn_idx_for_uuids(self, cn_idx: int, uuids: list[UUID]):
for uuid in uuids:
self.c_cn_idx.setdefault(uuid, []).append(cn_idx)
def _get_c_bank_for_uuids(self, uuids: list[UUID]): def _get_c_bank_for_uuids(self, uuids: list[UUID]):
per_i: list[list[Tensor]] = [] per_i: list[list[Tensor]] = []
for uuid in uuids: for uuid in uuids:
@@ -469,9 +457,10 @@ class BankStylesBasicTransformerBlock:
real_c_bank_list[i] = real_c_bank_list[i].to(cdevice) real_c_bank_list[i] = real_c_bank_list[i].to(cdevice)
return self.bank + real_c_bank_list return self.bank + real_c_bank_list
def _get_c_style_cfgs_for_uuids(self, uuids: list[UUID]):
# c_style_cfgs will be the same for all uuids def set_c_style_cfgs_for_uuids(self, style_cfg: float, uuids: list[UUID]):
return list(self.c_style_cfgs.values())[0] for uuid in uuids:
self.c_style_cfgs.setdefault(uuid, []).append(style_cfg)
def get_avg_style_fidelity(self, uuids: list[UUID], ignore_contextref): def get_avg_style_fidelity(self, uuids: list[UUID], ignore_contextref):
if ignore_contextref: if ignore_contextref:
@@ -479,15 +468,25 @@ class BankStylesBasicTransformerBlock:
combined = self.style_cfgs + self._get_c_style_cfgs_for_uuids(uuids) combined = self.style_cfgs + self._get_c_style_cfgs_for_uuids(uuids)
return sum(combined) / float(len(combined)) return sum(combined) / float(len(combined))
def _get_c_cn_idxs_for_uuids(self, uuids: list[UUID]): def _get_c_style_cfgs_for_uuids(self, uuids: list[UUID]):
# c_cn_idxs will be the same for all uids # c_style_cfgs will be the same for all provided uuids
return list(self.c_cn_idx.values())[0] return self.c_style_cfgs[uuids[0]]
def set_c_cn_idx_for_uuids(self, cn_idx: int, uuids: list[UUID]):
for uuid in uuids:
self.c_cn_idx.setdefault(uuid, []).append(cn_idx)
def get_cn_idxs(self, uuids: list[UUID], ignore_contxtref): def get_cn_idxs(self, uuids: list[UUID], ignore_contxtref):
if ignore_contxtref: if ignore_contxtref:
return self.cn_idx return self.cn_idx
return self.cn_idx + self._get_c_cn_idxs_for_uuids(uuids) return self.cn_idx + self._get_c_cn_idxs_for_uuids(uuids)
def _get_c_cn_idxs_for_uuids(self, uuids: list[UUID]):
# c_cn_idxs will be the same for all provided uuids
return self.c_cn_idx.get(uuids[0], [])
def init_cref_for_uuids(self, uuids: list[UUID]): def init_cref_for_uuids(self, uuids: list[UUID]):
for uuid in uuids: for uuid in uuids:
self.c_bank.setdefault(uuid, []) self.c_bank.setdefault(uuid, [])
@@ -529,48 +528,76 @@ class BankStylesTimestepEmbedSequential:
self.style_cfgs = [] self.style_cfgs = []
self.cn_idx: list[int] = [] self.cn_idx: list[int] = []
# cref # cref
self.c_var_bank: list[list] = [] self.c_var_bank: dict[UUID, list[Tensor]] = {}
self.c_mean_bank: list[list] = [] self.c_mean_bank: dict[UUID, list[Tensor]] = {}
self.c_style_cfgs: list[list] = [] self.c_style_cfgs: dict[UUID, list[float]] = {}
self.c_cn_idx: list[list[int]] = [] self.c_cn_idx: dict[UUID, list[int]] = {}
def get_var_bank(self, cref_idx, ignore_contextref): def set_c_var_bank_for_uuids(self, var: Tensor, uuids: list[UUID]):
if ignore_contextref or cref_idx >= len(self.c_var_bank): for uuid in uuids:
self.c_var_bank.setdefault(uuid, []).append(var)
def get_var_bank(self, uuids: list[UUID], ignore_contextref):
if ignore_contextref:
return self.var_bank return self.var_bank
return self.var_bank + self.c_var_bank[cref_idx] return self.var_bank + self._get_c_var_bank_for_uuids(uuids)
def _get_c_var_bank_for_uuids(self, uuids: list[UUID]):
return self.c_var_bank.get(uuids[0], [])
def get_mean_bank(self, cref_idx, ignore_contextref):
if ignore_contextref or cref_idx >= len(self.c_mean_bank): def set_c_mean_bank_for_uuids(self, mean: Tensor, uuids: list[UUID]):
for uuid in uuids:
self.c_mean_bank.setdefault(uuid, []).append(mean)
def get_mean_bank(self, uuids: list[UUID], ignore_contextref):
if ignore_contextref:
return self.mean_bank return self.mean_bank
return self.mean_bank + self.c_mean_bank[cref_idx] return self.mean_bank + self._get_c_mean_bank_for_uuids(uuids)
def get_style_cfgs(self, cref_idx, ignore_contextref): def _get_c_mean_bank_for_uuids(self, uuids: list[UUID]):
if ignore_contextref or cref_idx >= len(self.c_style_cfgs): return self.c_mean_bank.get(uuids[0], [])
def set_c_style_cfgs_for_uuids(self, style_cfg: float, uuids: list[UUID]):
for uuid in uuids:
self.c_style_cfgs.setdefault(uuid, []).append(style_cfg)
def get_style_cfgs(self, uuids: list[UUID], ignore_contextref):
if ignore_contextref:
return self.style_cfgs return self.style_cfgs
return self.style_cfgs + self.c_style_cfgs[cref_idx] return self.style_cfgs + self._get_c_style_cfgs_for_uuids(uuids)
def _get_c_style_cfgs_for_uuids(self, uuids: list[UUID]):
return self.c_style_cfgs.get(uuids[0], [])
def get_cn_idxs(self, cref_idx, ignore_contextref):
if ignore_contextref or cref_idx >= len(self.c_cn_idx): def set_c_cn_idx_for_uuids(self, cn_idx: int, uuids: list[UUID]):
for uuid in uuids:
self.c_cn_idx.setdefault(uuid, []).append(cn_idx)
def get_cn_idxs(self, uuids: list[UUID], ignore_contextref):
if ignore_contextref:
return self.cn_idx return self.cn_idx
return self.cn_idx + self.c_cn_idx[cref_idx] return self.cn_idx + self._get_c_cn_idxs_for_uuids(uuids)
def init_cref_for_idx(self, cref_idx: int): def _get_c_cn_idxs_for_uuids(self, uuids: list[UUID]):
# makes sure cref lists can accommodate cref_idx return self.c_cn_idx.get(uuids[0], [])
if cref_idx < 0:
return
while cref_idx >= len(self.c_var_bank):
self.c_var_bank.append([])
self.c_mean_bank.append([])
self.c_style_cfgs.append([])
self.c_cn_idx.append([])
def clear_cref_for_idx(self, cref_idx: int):
if cref_idx < 0 or cref_idx >= len(self.c_var_bank): def init_cref_for_uuids(self, uuids: list[UUID]):
return for uuid in uuids:
self.c_var_bank[cref_idx] = [] self.c_var_bank.setdefault(uuid, [])
self.c_mean_bank[cref_idx] = [] self.c_mean_bank.setdefault(uuid, [])
self.c_style_cfgs[cref_idx] = [] self.c_style_cfgs.setdefault(uuid, [])
self.c_cn_idx[cref_idx] = [] self.c_cn_idx.setdefault(uuid, [])
def clear_cref_for_uuids(self, uuids: list[UUID]):
for uuid in uuids:
self.c_var_bank[uuid] = []
self.c_mean_bank[uuid] = []
self.c_style_cfgs[uuid] = []
self.c_cn_idx[uuid] = []
def clean_ref(self): def clean_ref(self):
del self.mean_bank del self.mean_bank
@@ -587,10 +614,10 @@ class BankStylesTimestepEmbedSequential:
del self.c_mean_bank del self.c_mean_bank
del self.c_style_cfgs del self.c_style_cfgs
del self.c_cn_idx del self.c_cn_idx
self.c_var_bank = [] self.c_var_bank = {}
self.c_mean_bank = [] self.c_mean_bank = {}
self.c_style_cfgs = [] self.c_style_cfgs = {}
self.c_cn_idx = [] self.c_cn_idx = {}
def clean_all(self): def clean_all(self):
self.clean_ref() self.clean_ref()
@@ -872,9 +899,7 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten
# WRITE mode may have only 1 ReferenceAdvanced for RefCN at a time, other modes will have all ReferenceAdvanced # WRITE mode may have only 1 ReferenceAdvanced for RefCN at a time, other modes will have all ReferenceAdvanced
ref_write_cns: list[ReferenceAdvanced] = transformer_options.get(REF_WRITE_ATTN_CONTROL_LIST, []) ref_write_cns: list[ReferenceAdvanced] = transformer_options.get(REF_WRITE_ATTN_CONTROL_LIST, [])
ref_read_cns: list[ReferenceAdvanced] = transformer_options.get(REF_READ_ATTN_CONTROL_LIST, []) ref_read_cns: list[ReferenceAdvanced] = transformer_options.get(REF_READ_ATTN_CONTROL_LIST, [])
cref_cond_idx: int = transformer_options.get(CONTEXTREF_TEMP_COND_IDX, -1)
ignore_contextref_read = cref_mode in [MachineState.OFF, MachineState.WRITE] ignore_contextref_read = cref_mode in [MachineState.OFF, MachineState.WRITE]
#ignore_contextref_read = cref_cond_idx < 0 # if just writing to bank, should NOT be read in the same execution
#logger.info(f"cref: {cref_cond_idx}, cmode: {cref_mode}, ignored: {ignore_contextref_read}") #logger.info(f"cref: {cref_cond_idx}, cmode: {cref_mode}, ignored: {ignore_contextref_read}")
cached_n = None cached_n = None
@@ -1071,12 +1096,12 @@ def forward_timestep_embed_ref_inject_factory(orig_timestep_embed_inject_factory
# Reference CN stuff # Reference CN stuff
uc_idx_mask = transformer_options.get(REF_UNCOND_IDXS, []) uc_idx_mask = transformer_options.get(REF_UNCOND_IDXS, [])
uuids = transformer_options["uuids"] uuids = transformer_options["uuids"]
cref_mode = transformer_options.get(CONTEXTREF_MACHINE_STATE, MachineState.OFF)
#c_idx_mask = transformer_options.get(REF_COND_IDXS, []) #c_idx_mask = transformer_options.get(REF_COND_IDXS, [])
# WRITE mode will only have one ReferenceAdvanced, other modes will have all ReferenceAdvanced # WRITE mode will only have one ReferenceAdvanced, other modes will have all ReferenceAdvanced
ref_write_cns: list[ReferenceAdvanced] = transformer_options.get(REF_WRITE_ADAIN_CONTROL_LIST, []) ref_write_cns: list[ReferenceAdvanced] = transformer_options.get(REF_WRITE_ADAIN_CONTROL_LIST, [])
ref_read_cns: list[ReferenceAdvanced] = transformer_options.get(REF_READ_ADAIN_CONTROL_LIST, []) ref_read_cns: list[ReferenceAdvanced] = transformer_options.get(REF_READ_ADAIN_CONTROL_LIST, [])
cref_cond_idx: int = transformer_options.get(CONTEXTREF_TEMP_COND_IDX, -1) ignore_contextref_read = cref_mode in [MachineState.OFF, MachineState.WRITE]
ignore_contextref_read = cref_cond_idx < 0 # if writing to bank, should NOT be read in the same execution
cached_var = None cached_var = None
cached_mean = None cached_mean = None
@@ -1088,7 +1113,7 @@ def forward_timestep_embed_ref_inject_factory(orig_timestep_embed_inject_factory
cached_var, cached_mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0) cached_var, cached_mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0)
if refcn.is_context_ref: if refcn.is_context_ref:
cref_write_cns.append(refcn) cref_write_cns.append(refcn)
ts.injection_holder.bank_styles.init_cref_for_idx(cref_cond_idx) ts.injection_holder.bank_styles.init_cref_for_uuids(uuids)
else: else:
ts.injection_holder.bank_styles.var_bank.append(cached_var) ts.injection_holder.bank_styles.var_bank.append(cached_var)
ts.injection_holder.bank_styles.mean_bank.append(cached_mean) ts.injection_holder.bank_styles.mean_bank.append(cached_mean)
@@ -1100,16 +1125,16 @@ def forward_timestep_embed_ref_inject_factory(orig_timestep_embed_inject_factory
# if any refs to READ, do math with saved var, mean, and style_cfg # if any refs to READ, do math with saved var, mean, and style_cfg
if len(ref_read_cns) > 0: if len(ref_read_cns) > 0:
if len(ts.injection_holder.bank_styles.get_var_bank(cref_cond_idx, ignore_contextref_read)) > 0: if len(ts.injection_holder.bank_styles.get_cn_idxs(uuids, ignore_contextref_read)) > 0:
bank_styles = ts.injection_holder.bank_styles bank_styles = ts.injection_holder.bank_styles
var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0) var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0)
std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5
y_uc = torch.zeros_like(x) y_uc = torch.zeros_like(x)
cn_idx = 0 cn_idx = 0
real_style_cfgs = bank_styles.get_style_cfgs(cref_cond_idx, ignore_contextref_read) real_style_cfgs = bank_styles.get_style_cfgs(uuids, ignore_contextref_read)
real_var_bank = bank_styles.get_var_bank(cref_cond_idx, ignore_contextref_read) real_var_bank = bank_styles.get_var_bank(uuids, ignore_contextref_read)
real_mean_bank = bank_styles.get_mean_bank(cref_cond_idx, ignore_contextref_read) real_mean_bank = bank_styles.get_mean_bank(uuids, ignore_contextref_read)
real_cn_idxs = bank_styles.get_cn_idxs(cref_cond_idx, ignore_contextref_read) real_cn_idxs = bank_styles.get_cn_idxs(uuids, ignore_contextref_read)
for idx, order in enumerate(real_cn_idxs): for idx, order in enumerate(real_cn_idxs):
# make sure matching ref cn is selected # make sure matching ref cn is selected
for i in range(cn_idx, len(ref_read_cns)): for i in range(cn_idx, len(ref_read_cns)):
@@ -1138,13 +1163,13 @@ def forward_timestep_embed_ref_inject_factory(orig_timestep_embed_inject_factory
# ContextRef CN WRITE # ContextRef CN WRITE
if len(cref_write_cns) > 0: if len(cref_write_cns) > 0:
# clear so that ContextRef CNs can properly 'replace' previous value at cond_idx # clear so that ContextRef CNs can properly 'replace' previous value at cond_idx
ts.injection_holder.bank_styles.clear_cref_for_idx(cref_cond_idx) ts.injection_holder.bank_styles.clear_cref_for_uuids(uuids)
for refcn in cref_write_cns: for refcn in cref_write_cns:
# add a whole list to match expected type when combining # add a whole list to match expected type when combining
ts.injection_holder.bank_styles.c_var_bank[cref_cond_idx].append(cached_var) ts.injection_holder.bank_styles.set_c_var_bank_for_uuids(cached_var, uuids)
ts.injection_holder.bank_styles.c_mean_bank[cref_cond_idx].append(cached_mean) ts.injection_holder.bank_styles.set_c_mean_bank_for_uuids(cached_mean, uuids)
ts.injection_holder.bank_styles.c_style_cfgs[cref_cond_idx].append(refcn.ref_opts.adain_style_fidelity) ts.injection_holder.bank_styles.set_c_style_cfgs_for_uuids(refcn.ref_opts.adain_style_fidelity, uuids)
ts.injection_holder.bank_styles.c_cn_idx[cref_cond_idx].append(refcn.order) ts.injection_holder.bank_styles.set_c_cn_idx_for_uuids(refcn.order, uuids)
del cached_var del cached_var
del cached_mean del cached_mean