ctheodoris commited on
Commit
08914be
·
verified ·
1 Parent(s): c8dfd8d

revise to reinstate input_ids and move device check outside of loop

Browse files
Files changed (1) hide show
  1. geneformer/emb_extractor.py +5 -3
geneformer/emb_extractor.py CHANGED
@@ -92,13 +92,15 @@ def get_embs(
92
 
93
  overall_max_len = 0
94
 
 
 
95
  for i in trange(0, total_batch_length, forward_batch_size, leave=(not silent)):
96
  max_range = min(i + forward_batch_size, total_batch_length)
97
 
98
  minibatch = filtered_input_data.select([i for i in range(i, max_range)])
99
 
100
  max_len = int(max(minibatch["length"]))
101
- original_lens = torch.tensor(minibatch["length"], device="cuda" if torch.cuda.is_available() else "cpu")
102
  minibatch.set_format(type="torch")
103
 
104
  input_data_minibatch = minibatch["input_ids"]
@@ -108,8 +110,8 @@ def get_embs(
108
 
109
  with torch.no_grad():
110
  outputs = model(
111
- original_lens = torch.tensor(minibatch["length"], device="cuda" if torch.cuda.is_available() else "cpu")
112
- attention_mask=pu.gen_attention_mask(minibatch),
113
  )
114
 
115
  embs_i = outputs.hidden_states[layer_to_quant]
 
92
 
93
  overall_max_len = 0
94
 
95
+ device_type = "cuda" if torch.cuda.is_available() else "cpu"
96
+
97
  for i in trange(0, total_batch_length, forward_batch_size, leave=(not silent)):
98
  max_range = min(i + forward_batch_size, total_batch_length)
99
 
100
  minibatch = filtered_input_data.select([i for i in range(i, max_range)])
101
 
102
  max_len = int(max(minibatch["length"]))
103
+ original_lens = torch.tensor(minibatch["length"], device=device_type)
104
  minibatch.set_format(type="torch")
105
 
106
  input_data_minibatch = minibatch["input_ids"]
 
110
 
111
  with torch.no_grad():
112
  outputs = model(
113
+ input_ids=input_data_minibatch.to(device_type),
114
+ attention_mask=pu.gen_attention_mask(minibatch).to(device_type),
115
  )
116
 
117
  embs_i = outputs.hidden_states[layer_to_quant]