dlgenomics commited on
Commit
c8dfd8d
·
verified ·
1 Parent(s): 04c2b2e

Use dynamic CUDA check instead of hardcoded device in emb_extractor.py

Browse files

Fixes CPU-only inference, which currently fails with "AssertionError: Torch not compiled with CUDA enabled" because device="cuda" is hardcoded.

Replaces the hardcoded "cuda" with a dynamic check ("cuda" if torch.cuda.is_available() else "cpu"), matching the pattern already used elsewhere in the codebase (e.g. perturber_utils.py line 194).

See discussion #592: https://huggingface.co/ctheodoris/Geneformer/discussions/592

Files changed (1) hide show
  1. geneformer/emb_extractor.py +2 -2
geneformer/emb_extractor.py CHANGED
@@ -98,7 +98,7 @@ def get_embs(
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")
102
  minibatch.set_format(type="torch")
103
 
104
  input_data_minibatch = minibatch["input_ids"]
@@ -108,7 +108,7 @@ def get_embs(
108
 
109
  with torch.no_grad():
110
  outputs = model(
111
- input_ids=input_data_minibatch.to("cuda"),
112
  attention_mask=pu.gen_attention_mask(minibatch),
113
  )
114
 
 
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
 
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