From 76acae595dd49dcf39c5f97c3b57ec6a3010e3a8 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 1 Nov 2024 19:39:13 -0400 Subject: [PATCH] force vqa to use hfcache Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + modules/vqa.py | 24 ++++++++++++------------ 2 files changed, 13 insertions(+), 12 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 4dd30ed2c..9e5d830cc 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -37,6 +37,7 @@ This release can be considered an LTS release before we kick off the next round - fix diffusers load from folder - fix lora enum logging on windows - fix xyz grid with batch count + - fix vqa models ignoring hfcache folder setting - move downloads of some auxillary models to hfcache instead of models folder ## Update for 2024-10-29 diff --git a/modules/vqa.py b/modules/vqa.py index 7f0f8e17f..64ba83696 100644 --- a/modules/vqa.py +++ b/modules/vqa.py @@ -30,8 +30,8 @@ MODELS = { def git(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: - model = transformers.GitForCausalLM.from_pretrained(repo) - processor = transformers.GitProcessor.from_pretrained(repo) + model = transformers.GitForCausalLM.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) + processor = transformers.GitProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo model.to(devices.device, devices.dtype) shared.log.debug(f'VQA: class={model.__class__.__name__} processor={processor.__class__} model={repo}') @@ -55,8 +55,8 @@ def git(question: str, image: Image.Image, repo: str = None): def blip(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: - model = transformers.BlipForQuestionAnswering.from_pretrained(repo) - processor = transformers.BlipProcessor.from_pretrained(repo) + model = transformers.BlipForQuestionAnswering.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) + processor = transformers.BlipProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo model.to(devices.device, devices.dtype) inputs = processor(image, question, return_tensors="pt") @@ -73,8 +73,8 @@ def blip(question: str, image: Image.Image, repo: str = None): def vilt(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: - model = transformers.ViltForQuestionAnswering.from_pretrained(repo) - processor = transformers.ViltProcessor.from_pretrained(repo) + model = transformers.ViltForQuestionAnswering.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) + processor = transformers.ViltProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo model.to(devices.device) shared.log.debug(f'VQA: class={model.__class__.__name__} processor={processor.__class__} model={repo}') @@ -94,8 +94,8 @@ def vilt(question: str, image: Image.Image, repo: str = None): def pix(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: - model = transformers.Pix2StructForConditionalGeneration.from_pretrained(repo) - processor = transformers.Pix2StructProcessor.from_pretrained(repo) + model = transformers.Pix2StructForConditionalGeneration.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) + processor = transformers.Pix2StructProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo model.to(devices.device) shared.log.debug(f'VQA: class={model.__class__.__name__} processor={processor.__class__} model={repo}') @@ -115,8 +115,8 @@ def pix(question: str, image: Image.Image, repo: str = None): def moondream(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: - model = transformers.AutoModelForCausalLM.from_pretrained(repo, trust_remote_code=True) # revision = "2024-03-05" - processor = transformers.AutoTokenizer.from_pretrained(repo) # revision = "2024-03-05" + model = transformers.AutoModelForCausalLM.from_pretrained(repo, trust_remote_code=True, cache_dir=shared.opts.hfcache_dir) # revision = "2024-03-05" + processor = transformers.AutoTokenizer.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo model.eval() model.to(devices.device, devices.dtype) @@ -142,8 +142,8 @@ def florence(question: str, image: Image.Image, repo: str = None, revision: str return R if model is None or loaded != repo: transformers.dynamic_module_utils.get_imports = get_imports - model = transformers.AutoModelForCausalLM.from_pretrained(repo, trust_remote_code=True, revision=revision) - processor = transformers.AutoProcessor.from_pretrained(repo, trust_remote_code=True, revision=revision) + model = transformers.AutoModelForCausalLM.from_pretrained(repo, trust_remote_code=True, revision=revision, cache_dir=shared.opts.hfcache_dir) + processor = transformers.AutoProcessor.from_pretrained(repo, trust_remote_code=True, revision=revision, cache_dir=shared.opts.hfcache_dir) transformers.dynamic_module_utils.get_imports = _get_imports loaded = repo model.eval()