This commit is contained in:
silveroxides 2026-06-30 02:44:31 +02:00
parent 9fecaeb2d5
commit 0cfc1568b1

View File

@ -259,7 +259,7 @@ class SingleStreamDiT(nn.Module):
device = img.device device = img.device
ref_tokens_list, ref_pos_ids_list, ref_num_tokens = self._process_ref_latents( ref_tokens_list, ref_pos_ids_list, ref_num_tokens = self._process_ref_latents(
ref_latents, ref_latents_method, device, bs ref_latents, ref_latents_method, device, bs, h_, w_
) )
if len(ref_num_tokens) > 0: if len(ref_num_tokens) > 0:
@ -303,15 +303,15 @@ class SingleStreamDiT(nn.Module):
) )
return context.reshape(b, seq, self.txtlayers, self.txtdim) return context.reshape(b, seq, self.txtlayers, self.txtdim)
def _process_ref_latents(self, ref_latents, ref_latents_method, device, bs): def _process_ref_latents(self, ref_latents, ref_latents_method, device, bs, h_main, w_main):
ref_tokens_list = [] ref_tokens_list = []
ref_pos_ids_list = [] ref_pos_ids_list = []
ref_num_tokens = [] ref_num_tokens = []
patch = self.patch patch = self.patch
if ref_latents is not None: if ref_latents is not None:
h = 0 h = h_main
w = 0 w = w_main
index = 0 index = 0
index_ref_method = (ref_latents_method == "index") or (ref_latents_method == "index_timestep_zero") index_ref_method = (ref_latents_method == "index") or (ref_latents_method == "index_timestep_zero")
negative_ref_method = ref_latents_method == "negative_index" negative_ref_method = ref_latents_method == "negative_index"
@ -327,22 +327,19 @@ class SingleStreamDiT(nn.Module):
if index_ref_method: if index_ref_method:
index += 1 index += 1
gh_offset = 0
gw_offset = 0
elif negative_ref_method: elif negative_ref_method:
index -= 1 index -= 1
gh_offset = 0
gw_offset = 0
else: # offset/default else: # offset/default
index = 1 index = 0 # Stay in t_idx = 0 plane to remain 100% in-distribution
gh_offset = 0
gw_offset = 0 gh_offset = 0
if ref_gh + h > ref_gw + w: gw_offset = 0
gw_offset = w if ref_gh + h > ref_gw + w:
else: gw_offset = w
gh_offset = h else:
h = max(h, ref_gh + gh_offset) gh_offset = h
w = max(w, ref_gw + gw_offset) h = max(h, ref_gh + gh_offset)
w = max(w, ref_gw + gw_offset)
ref_tokens = rearrange(ref_pad, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch) ref_tokens = rearrange(ref_pad, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch)
ref_tokens = self.first(ref_tokens) ref_tokens = self.first(ref_tokens)