Ported ContextRef adain support to work with uuids, reorganized the ContextRef-related code in BankStyle classes
This commit is contained in:
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user