from dataclasses import replace from ppp import PromptPostProcessor from ppp_classes import ONWARNING_CHOICES # type: ignore from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") class TestModelVariants(TestPromptPostProcessorBase): def setUp(self): # pylint: disable=arguments-differ super().setUp(enable_file_logging=False) # Model variants tests def test_variants(self): self.process( InputTuple( "test1test2test3test4", "", ), OutputTuple("test1test2", ""), ppp=PromptPostProcessor( self.ppp_logger, { **self.def_env_info, "model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors", "ppp_config": { "models": { "sd1": { "detect": {"tests": {"class": ["SD15", "SD15_instructpix2pix"]}}, "variants": { "test3": {"find_in_filename": "testmodel"}, "sdxl": {"find_in_filename": "testmodel"}, }, }, "sdxl": { "detect": { "tests": { "class": [ "SDXL", "SDXLRefiner", "SDXL_instructpix2pix", "Segmind_Vega", "KOALA_700M", "KOALA_1B", ] } }, "variants": { "test1": {"find_in_filename": "testmodel"}, "test2": {"find_in_filename": "testmodel"}, }, }, "something": { "detect": {"tests": {"class": ["something"]}}, "variants": { "test4": {"find_in_filename": "testmodel"}, }, }, } }, }, replace( self.defopts, on_warning=ONWARNING_CHOICES.warn, ), self.grammar_content, self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), ) def test_variants_null_model(self): """null model in config disables detection and its variants""" self.process( InputTuple( "SDXLnot SDXL, PONYnot PONY", "", ), OutputTuple("not SDXL, not PONY", ""), ppp=PromptPostProcessor( self.ppp_logger, { **self.def_env_info, "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", "ppp_config": { "models": { "sdxl": None, } }, }, replace( self.defopts, on_warning=ONWARNING_CHOICES.warn, ), self.grammar_content, self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), )