MCPcopy Create free account
hub / github.com/TencentARC/BrushNet / load_notes_encoder

Function load_notes_encoder

scripts/convert_music_spectrogram_to_diffusers.py:19–43  ·  view source on GitHub ↗
(weights, model)

Source from the content-addressed store, hash-verified

17
18
19def load_notes_encoder(weights, model):
20 model.token_embedder.weight = nn.Parameter(torch.FloatTensor(weights["token_embedder"]["embedding"]))
21 model.position_encoding.weight = nn.Parameter(
22 torch.FloatTensor(weights["Embed_0"]["embedding"]), requires_grad=False
23 )
24 for lyr_num, lyr in enumerate(model.encoders):
25 ly_weight = weights[f"layers_{lyr_num}"]
26 lyr.layer[0].layer_norm.weight = nn.Parameter(
27 torch.FloatTensor(ly_weight["pre_attention_layer_norm"]["scale"])
28 )
29
30 attention_weights = ly_weight["attention"]
31 lyr.layer[0].SelfAttention.q.weight = nn.Parameter(torch.FloatTensor(attention_weights["query"]["kernel"].T))
32 lyr.layer[0].SelfAttention.k.weight = nn.Parameter(torch.FloatTensor(attention_weights["key"]["kernel"].T))
33 lyr.layer[0].SelfAttention.v.weight = nn.Parameter(torch.FloatTensor(attention_weights["value"]["kernel"].T))
34 lyr.layer[0].SelfAttention.o.weight = nn.Parameter(torch.FloatTensor(attention_weights["out"]["kernel"].T))
35
36 lyr.layer[1].layer_norm.weight = nn.Parameter(torch.FloatTensor(ly_weight["pre_mlp_layer_norm"]["scale"]))
37
38 lyr.layer[1].DenseReluDense.wi_0.weight = nn.Parameter(torch.FloatTensor(ly_weight["mlp"]["wi_0"]["kernel"].T))
39 lyr.layer[1].DenseReluDense.wi_1.weight = nn.Parameter(torch.FloatTensor(ly_weight["mlp"]["wi_1"]["kernel"].T))
40 lyr.layer[1].DenseReluDense.wo.weight = nn.Parameter(torch.FloatTensor(ly_weight["mlp"]["wo"]["kernel"].T))
41
42 model.layer_norm.weight = nn.Parameter(torch.FloatTensor(weights["encoder_norm"]["scale"]))
43 return model
44
45
46def load_continuous_encoder(weights, model):

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected