hugging-apps/echo-memory
0
1from .sd3_vae_encoder import SD3VAEEncoder, SDVAEEncoderStateDictConverter2from .sd3_vae_decoder import SD3VAEDecoder, SDVAEDecoderStateDictConverter3 4 5class FluxVAEEncoder(SD3VAEEncoder):6 def __init__(self):7 super().__init__()8 self.scaling_factor = 0.36119 self.shift_factor = 0.115910 11 @staticmethod12 def state_dict_converter():13 return FluxVAEEncoderStateDictConverter()14 15 16class FluxVAEDecoder(SD3VAEDecoder):17 def __init__(self):18 super().__init__()19 self.scaling_factor = 0.361120 self.shift_factor = 0.115921 22 @staticmethod23 def state_dict_converter():24 return FluxVAEDecoderStateDictConverter()25 26 27class FluxVAEEncoderStateDictConverter(SDVAEEncoderStateDictConverter):28 def __init__(self):29 pass30 31 def from_civitai(self, state_dict):32 rename_dict = {33 "encoder.conv_in.bias": "conv_in.bias",34 "encoder.conv_in.weight": "conv_in.weight",35 "encoder.conv_out.bias": "conv_out.bias",36 "encoder.conv_out.weight": "conv_out.weight",37 "encoder.down.0.block.0.conv1.bias": "blocks.0.conv1.bias",38 "encoder.down.0.block.0.conv1.weight": "blocks.0.conv1.weight",39 "encoder.down.0.block.0.conv2.bias": "blocks.0.conv2.bias",40 "encoder.down.0.block.0.conv2.weight": "blocks.0.conv2.weight",41 "encoder.down.0.block.0.norm1.bias": "blocks.0.norm1.bias",42 "encoder.down.0.block.0.norm1.weight": "blocks.0.norm1.weight",43 "encoder.down.0.block.0.norm2.bias": "blocks.0.norm2.bias",44 "encoder.down.0.block.0.norm2.weight": "blocks.0.norm2.weight",45 "encoder.down.0.block.1.conv1.bias": "blocks.1.conv1.bias",46 "encoder.down.0.block.1.conv1.weight": "blocks.1.conv1.weight",47 "encoder.down.0.block.1.conv2.bias": "blocks.1.conv2.bias",48 "encoder.down.0.block.1.conv2.weight": "blocks.1.conv2.weight",49 "encoder.down.0.block.1.norm1.bias": "blocks.1.norm1.bias",50 "encoder.down.0.block.1.norm1.weight": "blocks.1.norm1.weight",51 "encoder.down.0.block.1.norm2.bias": "blocks.1.norm2.bias",52 "encoder.down.0.block.1.norm2.weight": "blocks.1.norm2.weight",53 "encoder.down.0.downsample.conv.bias": "blocks.2.conv.bias",54 "encoder.down.0.downsample.conv.weight": "blocks.2.conv.weight",55 "encoder.down.1.block.0.conv1.bias": "blocks.3.conv1.bias",56 "encoder.down.1.block.0.conv1.weight": "blocks.3.conv1.weight",57 "encoder.down.1.block.0.conv2.bias": "blocks.3.conv2.bias",58 "encoder.down.1.block.0.conv2.weight": "blocks.3.conv2.weight",59 "encoder.down.1.block.0.nin_shortcut.bias": "blocks.3.conv_shortcut.bias",60 "encoder.down.1.block.0.nin_shortcut.weight": "blocks.3.conv_shortcut.weight",61 "encoder.down.1.block.0.norm1.bias": "blocks.3.norm1.bias",62 "encoder.down.1.block.0.norm1.weight": "blocks.3.norm1.weight",63 "encoder.down.1.block.0.norm2.bias": "blocks.3.norm2.bias",64 "encoder.down.1.block.0.norm2.weight": "blocks.3.norm2.weight",65 "encoder.down.1.block.1.conv1.bias": "blocks.4.conv1.bias",66 "encoder.down.1.block.1.conv1.weight": "blocks.4.conv1.weight",67 "encoder.down.1.block.1.conv2.bias": "blocks.4.conv2.bias",68 "encoder.down.1.block.1.conv2.weight": "blocks.4.conv2.weight",69 "encoder.down.1.block.1.norm1.bias": "blocks.4.norm1.bias",70 "encoder.down.1.block.1.norm1.weight": "blocks.4.norm1.weight",71 "encoder.down.1.block.1.norm2.bias": "blocks.4.norm2.bias",72 "encoder.down.1.block.1.norm2.weight": "blocks.4.norm2.weight",73 "encoder.down.1.downsample.conv.bias": "blocks.5.conv.bias",74 "encoder.down.1.downsample.conv.weight": "blocks.5.conv.weight",75 "encoder.down.2.block.0.conv1.bias": "blocks.6.conv1.bias",76 "encoder.down.2.block.0.conv1.weight": "blocks.6.conv1.weight",77 "encoder.down.2.block.0.conv2.bias": "blocks.6.conv2.bias",78 "encoder.down.2.block.0.conv2.weight": "blocks.6.conv2.weight",79 "encoder.down.2.block.0.nin_shortcut.bias": "blocks.6.conv_shortcut.bias",80 "encoder.down.2.block.0.nin_shortcut.weight": "blocks.6.conv_shortcut.weight",81 "encoder.down.2.block.0.norm1.bias": "blocks.6.norm1.bias",82 "encoder.down.2.block.0.norm1.weight": "blocks.6.norm1.weight",83 "encoder.down.2.block.0.norm2.bias": "blocks.6.norm2.bias",84 "encoder.down.2.block.0.norm2.weight": "blocks.6.norm2.weight",85 "encoder.down.2.block.1.conv1.bias": "blocks.7.conv1.bias",86 "encoder.down.2.block.1.conv1.weight": "blocks.7.conv1.weight",87 "encoder.down.2.block.1.conv2.bias": "blocks.7.conv2.bias",88 "encoder.down.2.block.1.conv2.weight": "blocks.7.conv2.weight",89 "encoder.down.2.block.1.norm1.bias": "blocks.7.norm1.bias",90 "encoder.down.2.block.1.norm1.weight": "blocks.7.norm1.weight",91 "encoder.down.2.block.1.norm2.bias": "blocks.7.norm2.bias",92 "encoder.down.2.block.1.norm2.weight": "blocks.7.norm2.weight",93 "encoder.down.2.downsample.conv.bias": "blocks.8.conv.bias",94 "encoder.down.2.downsample.conv.weight": "blocks.8.conv.weight",95 "encoder.down.3.block.0.conv1.bias": "blocks.9.conv1.bias",96 "encoder.down.3.block.0.conv1.weight": "blocks.9.conv1.weight",97 "encoder.down.3.block.0.conv2.bias": "blocks.9.conv2.bias",98 "encoder.down.3.block.0.conv2.weight": "blocks.9.conv2.weight",99 "encoder.down.3.block.0.norm1.bias": "blocks.9.norm1.bias",100 "encoder.down.3.block.0.norm1.weight": "blocks.9.norm1.weight",101 "encoder.down.3.block.0.norm2.bias": "blocks.9.norm2.bias",102 "encoder.down.3.block.0.norm2.weight": "blocks.9.norm2.weight",103 "encoder.down.3.block.1.conv1.bias": "blocks.10.conv1.bias",104 "encoder.down.3.block.1.conv1.weight": "blocks.10.conv1.weight",105 "encoder.down.3.block.1.conv2.bias": "blocks.10.conv2.bias",106 "encoder.down.3.block.1.conv2.weight": "blocks.10.conv2.weight",107 "encoder.down.3.block.1.norm1.bias": "blocks.10.norm1.bias",108 "encoder.down.3.block.1.norm1.weight": "blocks.10.norm1.weight",109 "encoder.down.3.block.1.norm2.bias": "blocks.10.norm2.bias",110 "encoder.down.3.block.1.norm2.weight": "blocks.10.norm2.weight",111 "encoder.mid.attn_1.k.bias": "blocks.12.transformer_blocks.0.to_k.bias",112 "encoder.mid.attn_1.k.weight": "blocks.12.transformer_blocks.0.to_k.weight",113 "encoder.mid.attn_1.norm.bias": "blocks.12.norm.bias",114 "encoder.mid.attn_1.norm.weight": "blocks.12.norm.weight",115 "encoder.mid.attn_1.proj_out.bias": "blocks.12.transformer_blocks.0.to_out.bias",116 "encoder.mid.attn_1.proj_out.weight": "blocks.12.transformer_blocks.0.to_out.weight",117 "encoder.mid.attn_1.q.bias": "blocks.12.transformer_blocks.0.to_q.bias",118 "encoder.mid.attn_1.q.weight": "blocks.12.transformer_blocks.0.to_q.weight",119 "encoder.mid.attn_1.v.bias": "blocks.12.transformer_blocks.0.to_v.bias",120 "encoder.mid.attn_1.v.weight": "blocks.12.transformer_blocks.0.to_v.weight",121 "encoder.mid.block_1.conv1.bias": "blocks.11.conv1.bias",122 "encoder.mid.block_1.conv1.weight": "blocks.11.conv1.weight",123 "encoder.mid.block_1.conv2.bias": "blocks.11.conv2.bias",124 "encoder.mid.block_1.conv2.weight": "blocks.11.conv2.weight",125 "encoder.mid.block_1.norm1.bias": "blocks.11.norm1.bias",126 "encoder.mid.block_1.norm1.weight": "blocks.11.norm1.weight",127 "encoder.mid.block_1.norm2.bias": "blocks.11.norm2.bias",128 "encoder.mid.block_1.norm2.weight": "blocks.11.norm2.weight",129 "encoder.mid.block_2.conv1.bias": "blocks.13.conv1.bias",130 "encoder.mid.block_2.conv1.weight": "blocks.13.conv1.weight",131 "encoder.mid.block_2.conv2.bias": "blocks.13.conv2.bias",132 "encoder.mid.block_2.conv2.weight": "blocks.13.conv2.weight",133 "encoder.mid.block_2.norm1.bias": "blocks.13.norm1.bias",134 "encoder.mid.block_2.norm1.weight": "blocks.13.norm1.weight",135 "encoder.mid.block_2.norm2.bias": "blocks.13.norm2.bias",136 "encoder.mid.block_2.norm2.weight": "blocks.13.norm2.weight",137 "encoder.norm_out.bias": "conv_norm_out.bias",138 "encoder.norm_out.weight": "conv_norm_out.weight",139 }140 state_dict_ = {}141 for name in state_dict:142 if name in rename_dict:143 param = state_dict[name]144 if "transformer_blocks" in rename_dict[name]:145 param = param.squeeze()146 state_dict_[rename_dict[name]] = param147 return state_dict_148 149 150 151class FluxVAEDecoderStateDictConverter(SDVAEDecoderStateDictConverter):152 def __init__(self):153 pass154 155 def from_civitai(self, state_dict):156 rename_dict = {157 "decoder.conv_in.bias": "conv_in.bias",158 "decoder.conv_in.weight": "conv_in.weight",159 "decoder.conv_out.bias": "conv_out.bias",160 "decoder.conv_out.weight": "conv_out.weight",161 "decoder.mid.attn_1.k.bias": "blocks.1.transformer_blocks.0.to_k.bias",162 "decoder.mid.attn_1.k.weight": "blocks.1.transformer_blocks.0.to_k.weight",163 "decoder.mid.attn_1.norm.bias": "blocks.1.norm.bias",164 "decoder.mid.attn_1.norm.weight": "blocks.1.norm.weight",165 "decoder.mid.attn_1.proj_out.bias": "blocks.1.transformer_blocks.0.to_out.bias",166 "decoder.mid.attn_1.proj_out.weight": "blocks.1.transformer_blocks.0.to_out.weight",167 "decoder.mid.attn_1.q.bias": "blocks.1.transformer_blocks.0.to_q.bias",168 "decoder.mid.attn_1.q.weight": "blocks.1.transformer_blocks.0.to_q.weight",169 "decoder.mid.attn_1.v.bias": "blocks.1.transformer_blocks.0.to_v.bias",170 "decoder.mid.attn_1.v.weight": "blocks.1.transformer_blocks.0.to_v.weight",171 "decoder.mid.block_1.conv1.bias": "blocks.0.conv1.bias",172 "decoder.mid.block_1.conv1.weight": "blocks.0.conv1.weight",173 "decoder.mid.block_1.conv2.bias": "blocks.0.conv2.bias",174 "decoder.mid.block_1.conv2.weight": "blocks.0.conv2.weight",175 "decoder.mid.block_1.norm1.bias": "blocks.0.norm1.bias",176 "decoder.mid.block_1.norm1.weight": "blocks.0.norm1.weight",177 "decoder.mid.block_1.norm2.bias": "blocks.0.norm2.bias",178 "decoder.mid.block_1.norm2.weight": "blocks.0.norm2.weight",179 "decoder.mid.block_2.conv1.bias": "blocks.2.conv1.bias",180 "decoder.mid.block_2.conv1.weight": "blocks.2.conv1.weight",181 "decoder.mid.block_2.conv2.bias": "blocks.2.conv2.bias",182 "decoder.mid.block_2.conv2.weight": "blocks.2.conv2.weight",183 "decoder.mid.block_2.norm1.bias": "blocks.2.norm1.bias",184 "decoder.mid.block_2.norm1.weight": "blocks.2.norm1.weight",185 "decoder.mid.block_2.norm2.bias": "blocks.2.norm2.bias",186 "decoder.mid.block_2.norm2.weight": "blocks.2.norm2.weight",187 "decoder.norm_out.bias": "conv_norm_out.bias",188 "decoder.norm_out.weight": "conv_norm_out.weight",189 "decoder.up.0.block.0.conv1.bias": "blocks.15.conv1.bias",190 "decoder.up.0.block.0.conv1.weight": "blocks.15.conv1.weight",191 "decoder.up.0.block.0.conv2.bias": "blocks.15.conv2.bias",192 "decoder.up.0.block.0.conv2.weight": "blocks.15.conv2.weight",193 "decoder.up.0.block.0.nin_shortcut.bias": "blocks.15.conv_shortcut.bias",194 "decoder.up.0.block.0.nin_shortcut.weight": "blocks.15.conv_shortcut.weight",195 "decoder.up.0.block.0.norm1.bias": "blocks.15.norm1.bias",196 "decoder.up.0.block.0.norm1.weight": "blocks.15.norm1.weight",197 "decoder.up.0.block.0.norm2.bias": "blocks.15.norm2.bias",198 "decoder.up.0.block.0.norm2.weight": "blocks.15.norm2.weight",199 "decoder.up.0.block.1.conv1.bias": "blocks.16.conv1.bias",200 "decoder.up.0.block.1.conv1.weight": "blocks.16.conv1.weight",201 "decoder.up.0.block.1.conv2.bias": "blocks.16.conv2.bias",202 "decoder.up.0.block.1.conv2.weight": "blocks.16.conv2.weight",203 "decoder.up.0.block.1.norm1.bias": "blocks.16.norm1.bias",204 "decoder.up.0.block.1.norm1.weight": "blocks.16.norm1.weight",205 "decoder.up.0.block.1.norm2.bias": "blocks.16.norm2.bias",206 "decoder.up.0.block.1.norm2.weight": "blocks.16.norm2.weight",207 "decoder.up.0.block.2.conv1.bias": "blocks.17.conv1.bias",208 "decoder.up.0.block.2.conv1.weight": "blocks.17.conv1.weight",209 "decoder.up.0.block.2.conv2.bias": "blocks.17.conv2.bias",210 "decoder.up.0.block.2.conv2.weight": "blocks.17.conv2.weight",211 "decoder.up.0.block.2.norm1.bias": "blocks.17.norm1.bias",212 "decoder.up.0.block.2.norm1.weight": "blocks.17.norm1.weight",213 "decoder.up.0.block.2.norm2.bias": "blocks.17.norm2.bias",214 "decoder.up.0.block.2.norm2.weight": "blocks.17.norm2.weight",215 "decoder.up.1.block.0.conv1.bias": "blocks.11.conv1.bias",216 "decoder.up.1.block.0.conv1.weight": "blocks.11.conv1.weight",217 "decoder.up.1.block.0.conv2.bias": "blocks.11.conv2.bias",218 "decoder.up.1.block.0.conv2.weight": "blocks.11.conv2.weight",219 "decoder.up.1.block.0.nin_shortcut.bias": "blocks.11.conv_shortcut.bias",220 "decoder.up.1.block.0.nin_shortcut.weight": "blocks.11.conv_shortcut.weight",221 "decoder.up.1.block.0.norm1.bias": "blocks.11.norm1.bias",222 "decoder.up.1.block.0.norm1.weight": "blocks.11.norm1.weight",223 "decoder.up.1.block.0.norm2.bias": "blocks.11.norm2.bias",224 "decoder.up.1.block.0.norm2.weight": "blocks.11.norm2.weight",225 "decoder.up.1.block.1.conv1.bias": "blocks.12.conv1.bias",226 "decoder.up.1.block.1.conv1.weight": "blocks.12.conv1.weight",227 "decoder.up.1.block.1.conv2.bias": "blocks.12.conv2.bias",228 "decoder.up.1.block.1.conv2.weight": "blocks.12.conv2.weight",229 "decoder.up.1.block.1.norm1.bias": "blocks.12.norm1.bias",230 "decoder.up.1.block.1.norm1.weight": "blocks.12.norm1.weight",231 "decoder.up.1.block.1.norm2.bias": "blocks.12.norm2.bias",232 "decoder.up.1.block.1.norm2.weight": "blocks.12.norm2.weight",233 "decoder.up.1.block.2.conv1.bias": "blocks.13.conv1.bias",234 "decoder.up.1.block.2.conv1.weight": "blocks.13.conv1.weight",235 "decoder.up.1.block.2.conv2.bias": "blocks.13.conv2.bias",236 "decoder.up.1.block.2.conv2.weight": "blocks.13.conv2.weight",237 "decoder.up.1.block.2.norm1.bias": "blocks.13.norm1.bias",238 "decoder.up.1.block.2.norm1.weight": "blocks.13.norm1.weight",239 "decoder.up.1.block.2.norm2.bias": "blocks.13.norm2.bias",240 "decoder.up.1.block.2.norm2.weight": "blocks.13.norm2.weight",241 "decoder.up.1.upsample.conv.bias": "blocks.14.conv.bias",242 "decoder.up.1.upsample.conv.weight": "blocks.14.conv.weight",243 "decoder.up.2.block.0.conv1.bias": "blocks.7.conv1.bias",244 "decoder.up.2.block.0.conv1.weight": "blocks.7.conv1.weight",245 "decoder.up.2.block.0.conv2.bias": "blocks.7.conv2.bias",246 "decoder.up.2.block.0.conv2.weight": "blocks.7.conv2.weight",247 "decoder.up.2.block.0.norm1.bias": "blocks.7.norm1.bias",248 "decoder.up.2.block.0.norm1.weight": "blocks.7.norm1.weight",249 "decoder.up.2.block.0.norm2.bias": "blocks.7.norm2.bias",250 "decoder.up.2.block.0.norm2.weight": "blocks.7.norm2.weight",251 "decoder.up.2.block.1.conv1.bias": "blocks.8.conv1.bias",252 "decoder.up.2.block.1.conv1.weight": "blocks.8.conv1.weight",253 "decoder.up.2.block.1.conv2.bias": "blocks.8.conv2.bias",254 "decoder.up.2.block.1.conv2.weight": "blocks.8.conv2.weight",255 "decoder.up.2.block.1.norm1.bias": "blocks.8.norm1.bias",256 "decoder.up.2.block.1.norm1.weight": "blocks.8.norm1.weight",257 "decoder.up.2.block.1.norm2.bias": "blocks.8.norm2.bias",258 "decoder.up.2.block.1.norm2.weight": "blocks.8.norm2.weight",259 "decoder.up.2.block.2.conv1.bias": "blocks.9.conv1.bias",260 "decoder.up.2.block.2.conv1.weight": "blocks.9.conv1.weight",261 "decoder.up.2.block.2.conv2.bias": "blocks.9.conv2.bias",262 "decoder.up.2.block.2.conv2.weight": "blocks.9.conv2.weight",263 "decoder.up.2.block.2.norm1.bias": "blocks.9.norm1.bias",264 "decoder.up.2.block.2.norm1.weight": "blocks.9.norm1.weight",265 "decoder.up.2.block.2.norm2.bias": "blocks.9.norm2.bias",266 "decoder.up.2.block.2.norm2.weight": "blocks.9.norm2.weight",267 "decoder.up.2.upsample.conv.bias": "blocks.10.conv.bias",268 "decoder.up.2.upsample.conv.weight": "blocks.10.conv.weight",269 "decoder.up.3.block.0.conv1.bias": "blocks.3.conv1.bias",270 "decoder.up.3.block.0.conv1.weight": "blocks.3.conv1.weight",271 "decoder.up.3.block.0.conv2.bias": "blocks.3.conv2.bias",272 "decoder.up.3.block.0.conv2.weight": "blocks.3.conv2.weight",273 "decoder.up.3.block.0.norm1.bias": "blocks.3.norm1.bias",274 "decoder.up.3.block.0.norm1.weight": "blocks.3.norm1.weight",275 "decoder.up.3.block.0.norm2.bias": "blocks.3.norm2.bias",276 "decoder.up.3.block.0.norm2.weight": "blocks.3.norm2.weight",277 "decoder.up.3.block.1.conv1.bias": "blocks.4.conv1.bias",278 "decoder.up.3.block.1.conv1.weight": "blocks.4.conv1.weight",279 "decoder.up.3.block.1.conv2.bias": "blocks.4.conv2.bias",280 "decoder.up.3.block.1.conv2.weight": "blocks.4.conv2.weight",281 "decoder.up.3.block.1.norm1.bias": "blocks.4.norm1.bias",282 "decoder.up.3.block.1.norm1.weight": "blocks.4.norm1.weight",283 "decoder.up.3.block.1.norm2.bias": "blocks.4.norm2.bias",284 "decoder.up.3.block.1.norm2.weight": "blocks.4.norm2.weight",285 "decoder.up.3.block.2.conv1.bias": "blocks.5.conv1.bias",286 "decoder.up.3.block.2.conv1.weight": "blocks.5.conv1.weight",287 "decoder.up.3.block.2.conv2.bias": "blocks.5.conv2.bias",288 "decoder.up.3.block.2.conv2.weight": "blocks.5.conv2.weight",289 "decoder.up.3.block.2.norm1.bias": "blocks.5.norm1.bias",290 "decoder.up.3.block.2.norm1.weight": "blocks.5.norm1.weight",291 "decoder.up.3.block.2.norm2.bias": "blocks.5.norm2.bias",292 "decoder.up.3.block.2.norm2.weight": "blocks.5.norm2.weight",293 "decoder.up.3.upsample.conv.bias": "blocks.6.conv.bias",294 "decoder.up.3.upsample.conv.weight": "blocks.6.conv.weight",295 }296 state_dict_ = {}297 for name in state_dict:298 if name in rename_dict:299 param = state_dict[name]300 if "transformer_blocks" in rename_dict[name]:301 param = param.squeeze()302 state_dict_[rename_dict[name]] = param303 return state_dict_