Add comments to ClipModel to notate changes from base

This commit is contained in:
Michael Poutre
2023-08-25 00:06:17 -07:00
parent 553f90691e
commit 0eaa2c1c41
+11 -1
View File
@@ -12,6 +12,7 @@ from custom_nodes.ClipStuff.lib.fun_clip_stuff import PromptLangTextModel
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
# Methods with no comment can be assumed to be the same as comfy.sd1_clip.SD1ClipModel
class SD1FunClipModel(torch.nn.Module):
"""Uses the CLIP transformer encoder for text (from huggingface)"""
LAYERS = [
@@ -27,14 +28,19 @@ class SD1FunClipModel(torch.nn.Module):
assert layer in self.LAYERS
self.num_layers = 12
if textmodel_path is not None:
# Our transformer
self.transformer = PromptLangTextModel.from_pretrained(textmodel_path)
else:
if textmodel_json_config is None:
# TODO: Maybe re-use clip config?
# Config could come from cond_stage_model.transformer.config
# Copied clip_config
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config.json")
config = CLIPTextConfig.from_json_file(textmodel_json_config)
self.num_layers = config.num_hidden_layers
with comfy.ops.use_comfy_ops():
with modeling_utils.no_init_weights():
# Our transformer
self.transformer = PromptLangTextModel(config)
self.max_length = max_length
@@ -68,6 +74,7 @@ class SD1FunClipModel(torch.nn.Module):
self.layer = self.layer_default[0]
self.layer_idx = self.layer_default[1]
# Completely changed to support Segments and actions
def set_up_textual_embeddings(self, tokens: list[list[SegOrAction]], current_embeds):
next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1
embedding_weights = []
@@ -131,6 +138,7 @@ class SD1FunClipModel(torch.nn.Module):
if segment.tokens[tokenIdx] == -1:
segment.tokens[tokenIdx] = n
# Support our set_up_textual_embeddings which modifies the input embeddings
def forward(self, tokens):
backup_embeds = self.transformer.get_input_embeddings()
device = backup_embeds.weight.device
@@ -166,7 +174,9 @@ class SD1FunClipModel(torch.nn.Module):
def load_sd(self, sd):
return self.transformer.load_state_dict(sd, strict=False)
def encode_token_weights(self, prompt_segments: list[list[SegOrAction]], **kwargs):
# Changed from comfy.sd1_clip.ClipTokenWeightEncoder
# Changed to use PromptSegments
def encode_token_weights(self, prompt_segments: list[list[SegOrAction]]):
to_encode = [[PromptSegment(text="_Empty Batch_", tokens=self.empty_tokens[0])]]
for batch in prompt_segments:
to_encode.append(batch)