From fd448e543726cb7a906740aa47ae98eca49745bf Mon Sep 17 00:00:00 2001 From: mayukhdeb Date: Thu, 20 Jun 2024 00:54:21 -0700 Subject: [PATCH] accomodate T5EncoderModel --- trainer/embedding_handler.py | 18 +++++++++++++++--- 1 file changed, 15 insertions(+), 3 deletions(-) diff --git a/trainer/embedding_handler.py b/trainer/embedding_handler.py index 4b81738..eac86f6 100644 --- a/trainer/embedding_handler.py +++ b/trainer/embedding_handler.py @@ -53,10 +53,22 @@ class TokenEmbeddingsHandler: continue # Ensure indices are a tensor. Use pre-existing dtype and device to match the model's. - indices_tensor = torch.tensor(indices, dtype=torch.long, device=text_encoder.text_model.embeddings.token_embedding.weight.device) + if isinstance(text_encoder, T5EncoderModel): + indices_tensor = torch.tensor( + indices, + dtype=torch.long, + device=text_encoder.encoder.embed_tokens.weight.device + ) + + # Directly access the embedding weights without detaching + token_embeddings = text_encoder.encoder.embed_tokens.weight[indices_tensor] + + else: + indices_tensor = torch.tensor(indices, dtype=torch.long, device=text_encoder.text_model.embeddings.token_embedding.weight.device) + + # Directly access the embedding weights without detaching + token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight[indices_tensor] - # Directly access the embedding weights without detaching - token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight[indices_tensor] embeddings[f'txt_encoder_{idx}'] = token_embeddings # Get all corresponding tokens for these embeddings