(lora, verbose, step="", emb_format=".pt")
| 49 | |
| 50 | |
| 51 | def unpack_bundle(lora, verbose, step="", emb_format=".pt"): |
| 52 | assert emb_format in [".pt", ".safetensors"] |
| 53 | if step != "": |
| 54 | step = "-" + str(step) |
| 55 | emb_dict = {} |
| 56 | bundle_keys = [] |
| 57 | for lora_key, value in lora.items(): |
| 58 | if lora_key.startswith("bundle_emb"): |
| 59 | bundle_keys.append(lora_key) |
| 60 | _, emb, *rest = lora_key.split(".") |
| 61 | emb = emb + step |
| 62 | if emb not in emb_dict: |
| 63 | emb_dict[emb] = {} |
| 64 | if len(rest) == 2: |
| 65 | key, subkey = rest |
| 66 | if emb_format == ".pt": |
| 67 | if key not in emb_dict[emb]: |
| 68 | emb_dict[emb][key] = {} |
| 69 | emb_dict[emb][key][subkey] = value |
| 70 | else: |
| 71 | emb_dict[emb][subkey] = value |
| 72 | elif len(rest) == 1: |
| 73 | key = rest[0] |
| 74 | emb_dict[emb][key] = value |
| 75 | for bundle_key in bundle_keys: |
| 76 | del lora[bundle_key] |
| 77 | if emb_format == ".pt": |
| 78 | for emb, emb_sd in emb_dict.items(): |
| 79 | emb_sd["name"] = emb |
| 80 | if verbose: |
| 81 | print("The following embeddings have been loaded from bundle") |
| 82 | print_emb_information(emb_dict) |
| 83 | return lora, emb_dict |
| 84 | |
| 85 | |
| 86 | def print_emb_information(emb_dict): |
no test coverage detected