make transform_DDP more intuitive

This commit is contained in:
ykume
2023-05-03 11:07:29 +09:00
parent e1143caf38
commit 2fcbfec178
6 changed files with 8 additions and 8 deletions

View File

@@ -315,7 +315,7 @@ def train(args):
)
# transform DDP after prepare
text_encoder, unet, _ = train_util.transform_DDP(text_encoder, unet)
text_encoder, unet = train_util.transform_if_model_is_DDP(text_encoder, unet)
index_no_updates = torch.arange(len(tokenizer)) < token_ids_XTI[0]
# print(len(index_no_updates), torch.sum(index_no_updates))