Skip to content

fix: allow LoRA prompts without explicit scale - #359

Open
andrewwhitecdw wants to merge 1 commit into
NVIDIA:mainfrom
andrewwhitecdw:bugfix/trt-allow-lora-prompts-without-explicit
Open

fix: allow LoRA prompts without explicit scale#359
andrewwhitecdw wants to merge 1 commit into
NVIDIA:mainfrom
andrewwhitecdw:bugfix/trt-allow-lora-prompts-without-explicit

Conversation

@andrewwhitecdw

Copy link
Copy Markdown

This PR addresses the following issue in scripts/trt.py: allow LoRA prompts without explicit scale.

Changes

  • scripts/trt.py: allow LoRA prompts without explicit scale.

Details

--- a/scripts/trt.py+++ b/scripts/trt.py@@ -1,7 +1,11 @@- # Get pathes- print("Apllying LoRAs: " + str(loras))- available = modelmanager.available_loras()- for lora in loras:- lora_name, lora_scale = lora.split(":")[1:]- lora_scales.append(float(lora_scale))- if lora_name not in available:+ # Get pathes+ print("Apllying LoRAs: " + str(loras))+ available = modelmanager.available_loras()+ for lora in loras:+ lora_parts = lora.split(":")[1:]+ if len(lora_parts) < 1:+ raise ValueError(f"Invalid LoRA tag: <{lora}>")+ lora_name = lora_parts[0]+ lora_scale = float(lora_parts[1]) if len(lora_parts) > 1 else 1.0+ lora_scales.append(lora_scale)+ if lora_name not in available:

Tests

  • tests/test_trt_lora.py
--- /dev/null+++ b/tests/test_trt_lora.py@@ -0,0 +1,56 @@+import unittest+from unittest.mock import MagicMock, patch++from scripts.trt import TensorRTScript+++class TestGetLoras(unittest.TestCase):+ def setUp(self):+ self.script = TensorRTScript()+ self.script.lora_hash = ""+ self.script.update_lora = False+ self.script.lora_refit_dict = {}++ def _prompt(self, text, model="model"):+ p = MagicMock()+ p.prompt = text+ p.sd_model_name = model+ return p++ @patch("scripts.trt.apply_loras")+ @patch("scripts.trt.modelmanager")+ def test_default_scale_for_lora_without_weight(self, modelmanager, apply_loras):+ modelmanager.available_loras.return_value = {"foo": "/path/foo.safetensors"}+ modelmanager.get_onnx_path.return_value = ("model", "/path/model.onnx")+ apply_loras.return_value = {}++ p = self._prompt("<lora:foo>")+ self.script.get_loras(p)++ self.assertEqual(p.prompt, "")+ apply_loras.assert_called_once_with(+ "/path/model.onnx",+ ["/path/foo.safetensors"],+ [1.0],+ )++ @patch("scripts.trt.apply_loras")+ @patch("scripts.trt.modelmanager")+ def test_explicit_scale_preserved(self, modelmanager, apply_loras):+ modelmanager.available_loras.return_value = {"bar": "/path/bar.safetensors"}+ modelmanager.get_onnx_path.return_value = ("model", "/path/model.onnx")+ apply_loras.return_value = {}++ p = self._prompt("<lora:bar:0.75>")+ self.script.get_loras(p)++ apply_loras.assert_called_once_with(+ "/path/model.onnx",+ ["/path/bar.safetensors"],+ [0.75],+ )+++if __name__ == "__main__":+ unittest.main()

Signed-off-by: andrewwhitecdw <andrewwhitecdw@users.noreply.github.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant

@andrewwhitecdw