I am writing a function with pytorch that feeds inputs through a transformer model, then condenses the last embedding layer by computing the average along a specific axis (using a subset of indices defined by a mask). Since the output of the model is very very large, I need to process the inputs in batches.
My question does not concern the logic of this function as I believe I have the correct implementation. My issue is that the function I wrote consumes memory excessively quickly and practically makes it unusable.
Here is my function:
def get_chunk_embeddings(encoded_dataset, batch_size):
chunk_embeddings = torch.empty([0,768])
for i in range(len(encoded_dataset['input_ids'])//batch_size):
input_ids = encoded_dataset['input_ids'][i*batch_size:i*batch_size + batch_size]
attention_mask = encoded_dataset['attention_mask'][i*batch_size:i*batch_size + batch_size]
embeddings = model.forward(input_ids=input_ids, attention_mask=attention_mask)['last_hidden_state']
embeddings = embeddings * attention_mask[:,:,None]
embeddings = embeddings.sum(dim=1)/attention_mask.sum(dim=1)[:,None]
chunk_embeddings = torch.cat([chunk_embeddings, embeddings],0)
return chunk_embeddings
Now let's talk memory (the numbers below assume that I pass a batch_size of 8):
- I am using google colab and I have ~25 GB of RAM available
- The
modelis a BERT model and consumes 413 MB encoded_datasetconsumes 0.48 GBinput_idsconsumes 0.413 MBattention_maskconsumes 4.096 KBembeddingsat its peak consumption consume 12.6 MBchunk_embeddingsadds 0.024576 MB with each iteration
So from my understanding, I should be able to allow chunk_embeddings to grow up to: 25GB - 413MB - 0.48GB - 0.413MB - 4.096KB - 12.6MB ~= 24 GB. Enough for almost 1 million iterations..
Here I will walk through an example of what I am experiencing:
- Before running my function, google colab tells me that I have plenty of memory
Now, for the sake of example, I will run the function (for only 3 iterations) To be explicit, I put this at the end of my for loop:
if (i == 2):return chunk_embeddingsNow I run the code
val = get_chunk_embeddings(train_encoded_dataset, 8)So even with just 3 iterations, somehow I consume almost 5.5 GB of RAM.
Why is this happening? Also after I have returned from the function all of the local variables should be deleted and there is no way that val is so large.
Could someone tell me what it is I am doing wrong or am not understanding? Please let me know if further information is needed.

