(func)
| 19 | return wrapper |
| 20 | |
| 21 | def copied_cache_model(func): |
| 22 | def wrapper(*args, **kwargs): |
| 23 | if ENABLE_CPU_CACHE: |
| 24 | model_name = func.__name__ + str(args) + str(kwargs) |
| 25 | if model_name not in cached_models: |
| 26 | cached_models[model_name] = func(*args, **kwargs) |
| 27 | return deepcopy(cached_models[model_name]) |
| 28 | else: |
| 29 | return func(*args, **kwargs) |
| 30 | return wrapper |
| 31 | |
| 32 | def model_from_ckpt_or_pretrained(ckpt_or_pretrained, model_cls, original_config_file='ckpt/v1-inference.yaml', torch_dtype=torch.float16, **kwargs): |
| 33 | if ckpt_or_pretrained.endswith(".safetensors"): |
nothing calls this directly
no outgoing calls
no test coverage detected