diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 822576a..e58ec6d 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -426,23 +426,11 @@ class BankStylesBasicTransformerBlock: self.c_style_cfgs: dict[UUID, list[float]] = {} 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]): per_uuid = len(x) // len(uuids) for uuid, i in zip(uuids, list(range(0, len(x), 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]): per_i: list[list[Tensor]] = [] for uuid in uuids: @@ -469,9 +457,10 @@ class BankStylesBasicTransformerBlock: real_c_bank_list[i] = real_c_bank_list[i].to(cdevice) 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 - return list(self.c_style_cfgs.values())[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_avg_style_fidelity(self, uuids: list[UUID], ignore_contextref): if ignore_contextref: @@ -479,15 +468,25 @@ class BankStylesBasicTransformerBlock: combined = self.style_cfgs + self._get_c_style_cfgs_for_uuids(uuids) return sum(combined) / float(len(combined)) - def _get_c_cn_idxs_for_uuids(self, uuids: list[UUID]): - # c_cn_idxs will be the same for all uids - return list(self.c_cn_idx.values())[0] + def _get_c_style_cfgs_for_uuids(self, uuids: list[UUID]): + # c_style_cfgs will be the same for all provided uuids + 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): if ignore_contxtref: return self.cn_idx 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]): for uuid in uuids: self.c_bank.setdefault(uuid, []) @@ -529,48 +528,76 @@ class BankStylesTimestepEmbedSequential: self.style_cfgs = [] self.cn_idx: list[int] = [] # cref - self.c_var_bank: list[list] = [] - self.c_mean_bank: list[list] = [] - self.c_style_cfgs: list[list] = [] - self.c_cn_idx: list[list[int]] = [] + self.c_var_bank: dict[UUID, list[Tensor]] = {} + self.c_mean_bank: dict[UUID, list[Tensor]] = {} + self.c_style_cfgs: dict[UUID, list[float]] = {} + self.c_cn_idx: dict[UUID, list[int]] = {} - def get_var_bank(self, cref_idx, ignore_contextref): - if ignore_contextref or cref_idx >= len(self.c_var_bank): + def set_c_var_bank_for_uuids(self, var: Tensor, uuids: list[UUID]): + 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 + 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 + 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): - if ignore_contextref or cref_idx >= len(self.c_style_cfgs): + def _get_c_mean_bank_for_uuids(self, uuids: list[UUID]): + 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 + 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 + 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): - # makes sure cref lists can accommodate cref_idx - 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 _get_c_cn_idxs_for_uuids(self, uuids: list[UUID]): + return self.c_cn_idx.get(uuids[0], []) - def clear_cref_for_idx(self, cref_idx: int): - if cref_idx < 0 or cref_idx >= len(self.c_var_bank): - return - self.c_var_bank[cref_idx] = [] - self.c_mean_bank[cref_idx] = [] - self.c_style_cfgs[cref_idx] = [] - self.c_cn_idx[cref_idx] = [] + + def init_cref_for_uuids(self, uuids: list[UUID]): + for uuid in uuids: + self.c_var_bank.setdefault(uuid, []) + self.c_mean_bank.setdefault(uuid, []) + self.c_style_cfgs.setdefault(uuid, []) + 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): del self.mean_bank @@ -587,10 +614,10 @@ class BankStylesTimestepEmbedSequential: del self.c_mean_bank del self.c_style_cfgs del self.c_cn_idx - self.c_var_bank = [] - self.c_mean_bank = [] - self.c_style_cfgs = [] - self.c_cn_idx = [] + self.c_var_bank = {} + self.c_mean_bank = {} + self.c_style_cfgs = {} + self.c_cn_idx = {} def clean_all(self): 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 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, []) - 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 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}") cached_n = None @@ -1071,12 +1096,12 @@ def forward_timestep_embed_ref_inject_factory(orig_timestep_embed_inject_factory # Reference CN stuff uc_idx_mask = transformer_options.get(REF_UNCOND_IDXS, []) uuids = transformer_options["uuids"] + cref_mode = transformer_options.get(CONTEXTREF_MACHINE_STATE, MachineState.OFF) #c_idx_mask = transformer_options.get(REF_COND_IDXS, []) # 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_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_cond_idx < 0 # if writing to bank, should NOT be read in the same execution + ignore_contextref_read = cref_mode in [MachineState.OFF, MachineState.WRITE] cached_var = 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) if refcn.is_context_ref: 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: ts.injection_holder.bank_styles.var_bank.append(cached_var) 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 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 var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0) std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 y_uc = torch.zeros_like(x) cn_idx = 0 - real_style_cfgs = bank_styles.get_style_cfgs(cref_cond_idx, ignore_contextref_read) - real_var_bank = bank_styles.get_var_bank(cref_cond_idx, ignore_contextref_read) - real_mean_bank = bank_styles.get_mean_bank(cref_cond_idx, ignore_contextref_read) - real_cn_idxs = bank_styles.get_cn_idxs(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(uuids, ignore_contextref_read) + real_mean_bank = bank_styles.get_mean_bank(uuids, ignore_contextref_read) + real_cn_idxs = bank_styles.get_cn_idxs(uuids, ignore_contextref_read) for idx, order in enumerate(real_cn_idxs): # make sure matching ref cn is selected 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 if len(cref_write_cns) > 0: # 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: # 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.c_mean_bank[cref_cond_idx].append(cached_mean) - ts.injection_holder.bank_styles.c_style_cfgs[cref_cond_idx].append(refcn.ref_opts.adain_style_fidelity) - ts.injection_holder.bank_styles.c_cn_idx[cref_cond_idx].append(refcn.order) + ts.injection_holder.bank_styles.set_c_var_bank_for_uuids(cached_var, uuids) + ts.injection_holder.bank_styles.set_c_mean_bank_for_uuids(cached_mean, uuids) + ts.injection_holder.bank_styles.set_c_style_cfgs_for_uuids(refcn.ref_opts.adain_style_fidelity, uuids) + ts.injection_holder.bank_styles.set_c_cn_idx_for_uuids(refcn.order, uuids) del cached_var del cached_mean