From 2fb95718e4529727383c48e5cb34ddb6275d82a7 Mon Sep 17 00:00:00 2001 From: haden Date: Thu, 25 Jan 2024 04:32:48 -0800 Subject: [PATCH] add cuda support --- moondream/text_model.py | 8 ++++++-- moondream/vision_encoder.py | 6 ++++-- 2 files changed, 10 insertions(+), 4 deletions(-) diff --git a/moondream/text_model.py b/moondream/text_model.py index 321b4b8..90b5168 100644 --- a/moondream/text_model.py +++ b/moondream/text_model.py @@ -12,6 +12,10 @@ transformers.logging.set_verbosity_error() class TextModel: def __init__(self, model_path: str = "model") -> None: super().__init__() + + # Determine if CUDA (GPU) is available and use it; otherwise, use CPU + self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + self.tokenizer = Tokenizer.from_pretrained(f"{model_path}/tokenizer") phi_config = PhiConfig.from_pretrained(f"{model_path}/text_model_cfg.json") @@ -21,10 +25,10 @@ class TextModel: self.model = load_checkpoint_and_dispatch( self.model, f"{model_path}/text_model.pt", - device_map={"": "cpu"}, + device_map={"": self.device.type}, ) - self.text_emb = self.model.get_input_embeddings() + self.text_emb = self.model.get_input_embeddings().to(self.device) def input_embeds(self, prompt, image_embeds): embeds = [] diff --git a/moondream/vision_encoder.py b/moondream/vision_encoder.py index 2122209..68848ca 100644 --- a/moondream/vision_encoder.py +++ b/moondream/vision_encoder.py @@ -13,7 +13,9 @@ from torchvision.transforms.v2 import ( class VisionEncoder: def __init__(self, model_path: str = "model") -> None: - self.model = torch.jit.load(f"{model_path}/vision.pt").to(dtype=torch.float32) + # Determine if CUDA (GPU) is available and use it; otherwise, use CPU + self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + self.model = torch.jit.load(f"{model_path}/vision.pt").to(self.device).to(dtype=torch.float32) self.preprocess = Compose( [ Resize(size=(384, 384), interpolation=InterpolationMode.BICUBIC), @@ -25,7 +27,7 @@ class VisionEncoder: def __call__(self, image: Image) -> torch.Tensor: with torch.no_grad(): - image_vec = self.preprocess(image.convert("RGB")).unsqueeze(0) + image_vec = self.preprocess(image.convert("RGB")).unsqueeze(0).to(self.device) image_vec = image_vec[:, :, :-6, :-6] image_vec = rearrange( image_vec, "b c (h p1) (w p2) -> b (h w) (c p1 p2)", p1=14, p2=14