support flux pulid

This commit is contained in:
TangYanxin
2025-01-14 22:41:36 +08:00
parent 0af82b5d9d
commit a0750c4e6f
+23
View File
@@ -57,6 +57,7 @@ def teacache_flux_forward(
img_mod1, _ = self.double_blocks[0].img_mod(vec_)
modulated_inp = self.double_blocks[0].img_norm1(inp)
modulated_inp = (1 + img_mod1.scale) * modulated_inp + img_mod1.shift
ca_idx = 0
if not hasattr(self, 'accumulated_rel_l1_distance'):
should_calc = True
@@ -113,6 +114,16 @@ def teacache_flux_forward(
if add is not None:
img += add
# PuLID attention
if getattr(self, "pulid_data", {}):
if i % self.pulid_double_interval == 0:
# Will calculate influence of all pulid nodes at once
for _, node_data in self.pulid_data.items():
if torch.any((node_data['sigma_start'] >= timesteps)
& (timesteps >= node_data['sigma_end'])):
img = img + node_data['weight'] * self.pulid_ca[ca_idx](node_data['embedding'], img)
ca_idx += 1
img = torch.cat((txt, img), 1)
for i, block in enumerate(self.single_blocks):
@@ -141,6 +152,18 @@ def teacache_flux_forward(
if add is not None:
img[:, txt.shape[1] :, ...] += add
# PuLID attention
if getattr(self, "pulid_data", {}):
real_img, txt = img[:, txt.shape[1]:, ...], img[:, :txt.shape[1], ...]
if i % self.pulid_single_interval == 0:
# Will calculate influence of all nodes at once
for _, node_data in self.pulid_data.items():
if torch.any((node_data['sigma_start'] >= timesteps)
& (timesteps >= node_data['sigma_end'])):
real_img = real_img + node_data['weight'] * self.pulid_ca[ca_idx](node_data['embedding'], real_img)
ca_idx += 1
img = torch.cat((txt, real_img), 1)
img = img[:, txt.shape[1] :, ...]
self.previous_residual = img - ori_img