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

#593

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

Thank you - I think input_ids was inadvertently overwritten in the forward pass so I just reinstated that and moved the device check to be outside the loop so it only does it once. Please test to make sure this works and let us know if any other issues arise. Thank you for your contribution to the codebase!

ctheodoris changed pull request status to closed

Sign up or log in to comment