MCPcopy Create free account
hub / github.com/TencentARC/BrushNet / save_new_embed

Function save_new_embed

examples/custom_diffusion/train_custom_diffusion.py:310–322  ·  view source on GitHub ↗

Saves the new token embeddings from the text encoder.

(text_encoder, modifier_token_id, accelerator, args, output_dir, safe_serialization=True)

Source from the content-addressed store, hash-verified

308
309
310def save_new_embed(text_encoder, modifier_token_id, accelerator, args, output_dir, safe_serialization=True):
311 """Saves the new token embeddings from the text encoder."""
312 logger.info("Saving embeddings")
313 learned_embeds = accelerator.unwrap_model(text_encoder).get_input_embeddings().weight
314 for x, y in zip(modifier_token_id, args.modifier_token):
315 learned_embeds_dict = {}
316 learned_embeds_dict[y] = learned_embeds[x]
317 filename = f"{output_dir}/{y}.bin"
318
319 if safe_serialization:
320 safetensors.torch.save_file(learned_embeds_dict, filename, metadata={"format": "pt"})
321 else:
322 torch.save(learned_embeds_dict, filename)
323
324
325def parse_args(input_args=None):

Callers 1

mainFunction · 0.85

Calls 2

infoMethod · 0.80
get_input_embeddingsMethod · 0.45

Tested by

no test coverage detected