(self, x, context=None, mask=None, pe=None, transformer_options={})
| 16 | |
| 17 | class LTXModifiedCrossAttention(nn.Module): |
| 18 | def forward(self, x, context=None, mask=None, pe=None, transformer_options={}): |
| 19 | context = x if context is None else context |
| 20 | context_v = x if context is None else context |
| 21 | |
| 22 | step = transformer_options.get("step", -1) |
| 23 | total_steps = transformer_options.get("total_steps", 0) |
| 24 | attn_bank = transformer_options.get("attn_bank", None) |
| 25 | sample_mode = transformer_options.get("sample_mode", None) |
| 26 | if attn_bank is not None and self.idx in attn_bank["block_map"]: |
| 27 | len_conds = len(transformer_options["cond_or_uncond"]) |
| 28 | pred_order = transformer_options["pred_order"] |
| 29 | if ( |
| 30 | sample_mode == "forward" |
| 31 | and total_steps - step - 1 < attn_bank["save_steps"] |
| 32 | ): |
| 33 | step_idx = f"{pred_order}_{total_steps-step-1}" |
| 34 | attn_bank["block_map"][self.idx][step_idx] = x.cpu() |
| 35 | elif sample_mode == "reverse" and step < attn_bank["inject_steps"]: |
| 36 | step_idx = f"{pred_order}_{step}" |
| 37 | inject_settings = attn_bank.get("inject_settings", {}) |
| 38 | if len(inject_settings) > 0: |
| 39 | inj = ( |
| 40 | attn_bank["block_map"][self.idx][step_idx] |
| 41 | .to(x.device) |
| 42 | .repeat(len_conds, 1, 1) |
| 43 | ) |
| 44 | if "q" in inject_settings: |
| 45 | x = inj |
| 46 | if "k" in inject_settings: |
| 47 | context = inj |
| 48 | if "v" in inject_settings: |
| 49 | context_v = inj |
| 50 | |
| 51 | q = self.to_q(x) |
| 52 | k = self.to_k(context) |
| 53 | v = self.to_v(context_v) |
| 54 | |
| 55 | q = self.q_norm(q) |
| 56 | k = self.k_norm(k) |
| 57 | |
| 58 | if pe is not None: |
| 59 | q = apply_rotary_emb(q, pe) |
| 60 | k = apply_rotary_emb(k, pe) |
| 61 | |
| 62 | feta_score = None |
| 63 | if ( |
| 64 | transformer_options.get("feta_weight", 0) > 0 |
| 65 | and self.idx in transformer_options["feta_layers"]["layers"] |
| 66 | ): |
| 67 | feta_score = get_feta_scores(q, k, self.heads, transformer_options) |
| 68 | |
| 69 | alt_attn_fn = ( |
| 70 | transformer_options.get("patches_replace", {}) |
| 71 | .get("layer", {}) |
| 72 | .get(("self_attn", self.idx), None) |
| 73 | ) |
| 74 | if alt_attn_fn is not None: |
| 75 | out = alt_attn_fn( |
nothing calls this directly
no test coverage detected