Vansh Chugh
remove dead code found by find_dead_methods.py
554d1bc
Raw
History Blame Contribute Delete
15.5 kB
import torch
from torch import nn
from .musicgen_cc import MusicGen
from .utilities.model_utils import freeze, print_trainable_parameters
import peft
def get_musicgen(sec, device):
mg = MusicGen.get_pretrained(name='large', device=device)
# org repo sets extend_stride=16 which causes huge offsets
# (extend_stride -> stride_tokens -> current_gen_offset)
# sec - 1 keeps a 1s overlap between chunks for continuity.
mg.set_generation_params(duration=sec, extend_stride=sec - 1, top_k=250)
mg.lm.here()
freeze(mg.lm)
return mg
class CondMusicgen(nn.Module):
def __init__(self, sec, device="cuda", top_k=250):
super().__init__()
mg = get_musicgen(sec, device)
mg.generation_params["top_k"] = top_k
self.musicgen = mg
self.lm = mg.lm
self.max_duration = sec
self.frame_rate = 50
def set_training(self):
self.lm.train()
def forward(self, input_code, text_description, embed_fn, num_samples=1, mode="train",
total_gen_len=None, prompt_tokens=None):
mg = self.musicgen
lm = self.lm
# attributes, _ = mg._prepare_tokens_and_attributes(text_description, None)
if mode == "train":
with mg.autocast:
out = lm.compute_predictions(codes=input_code,
embed_fn=embed_fn,
conditions=text_description)
return out
elif mode == "inference":
if total_gen_len is None:
total_gen_len = int(mg.duration * mg.frame_rate)
with mg.autocast:
gen_tokens = lm.generate(embed_fn=embed_fn, num_samples=num_samples,
prompt=None, conditions=text_description,
callback=None, max_gen_len=total_gen_len,
**mg.generation_params)
return gen_tokens
elif mode == "continuation":
with mg.autocast:
#if prompt_tokens is not None:
# print(prompt_tokens.shape)
gen_tokens = lm.generate(embed_fn=embed_fn, num_samples=num_samples,
prompt=prompt_tokens, conditions=text_description,
callback=None, max_gen_len=total_gen_len, **mg.generation_params)
return gen_tokens
def generate(self, cp_fn,
text_description,
condition_audio_code,
num_samples):
mg = self.musicgen
lm = self.lm
attributes, _ = mg._prepare_tokens_and_attributes(text_description, None)
all_tokens = []
stride_tokens = int(self.frame_rate * mg.extend_stride)
current_gen_offset = 0
prompt_length = 0
prompt_tokens = None
total_gen_len = condition_audio_code.shape[-1] - 1
total_sec = total_gen_len / 50.
while current_gen_offset + prompt_length < total_gen_len:
time_offset = current_gen_offset / self.frame_rate
chunk_duration = min(total_sec - time_offset, self.max_duration)
max_gen_len = int(chunk_duration * self.frame_rate)
if prompt_length >= max_gen_len:
break
#print("current_gen_offset / total ", current_gen_offset, "/", total_gen_len)
with mg.autocast:
condition_audio_code_clip = condition_audio_code[:, :, current_gen_offset:current_gen_offset + max_gen_len + 1]
#print(cond_mask.shape, drums_clip.shape, piano_roll_clip.shape, chords_clip.shape, max_gen_len)
embed_fn = cp_fn(condition_audio_code=condition_audio_code_clip,
max_n_frames=max_gen_len,
mode="inference")
gen_tokens = lm.generate(num_samples=num_samples,
embed_fn=embed_fn,
prompt=prompt_tokens,
conditions=attributes,
callback=None, max_gen_len=max_gen_len, **mg.generation_params)
if prompt_tokens is None:
all_tokens.append(gen_tokens)
else:
all_tokens.append(gen_tokens[:, :, prompt_tokens.shape[-1]:])
prompt_tokens = gen_tokens[:, :, stride_tokens:]
prompt_length = prompt_tokens.shape[-1]
current_gen_offset += stride_tokens
if current_gen_offset > 50 * 80:
break
gen_tokens = torch.cat(all_tokens, dim=-1)
return gen_tokens
class EmbFn:
def __init__(self, activates, fn, start_layer, max_len, inference=False, skip=None):
self.interval = None
self.index = -1
self.adaptor = None
self.start_layer = start_layer
self.activates = activates
self.max_len = max_len
self.fn = fn
self.inference = inference
self.skip = skip
def get_adaptor(self, tag):
index = self.index
if index < self.start_layer or tag == "cross":
return None, None
i = index - self.start_layer
adaptor, gate = self.fn(i, self.activates)
# if self.adaptor is not None:
# adaptor = self.adaptor + adaptor
return adaptor, gate
def clear_state(self):
self.qkv = {}
torch.cuda.empty_cache()
def crop(self, tag, x):
if self.interval is not None:
st, ed = self.interval
if st >= self.max_len:
st = self.max_len - 1
ed = st + 1
return x[:, :, st:ed, :]
return x
def get_cross_attention_src(self, src):
return src
def modify(self, x, dt_x, gate):
# return dt_x * gate[:dt_x.shape[-2], :] + x
return dt_x * gate + x
def update_adaptor(self, adaptor):
self.adaptor = adaptor
def set_index(self, index):
self.index = index
def update_interval(self, st, ed):
self.interval = [st, ed]
class CPTransformerLayer(nn.Module):
def __init__(self, norm1, norm2, layer_scale_1, dropout1, self_attn, layer_scale_2,
autocast, linear1, linear2, activation, dropout, *args, **kwargs):
super().__init__(*args, **kwargs)
self.norm1 = norm1
self.norm2 = norm2
self.layer_scale_1 = layer_scale_1
self.dropout1 = dropout1
self.self_attn = self_attn
self.layer_scale_2 = layer_scale_2
self.autocast = autocast
self.linear1 = linear1
self.linear2 = linear2
self.activation = activation
self.dropout = dropout
def _ff_block(self, x):
return self.linear2(self.dropout(self.activation(self.linear1(x))))
def forward(self, x, cond=None):
with self.autocast:
if cond is None:
nx = self.norm1(x)
else:
nx = self.norm1(x) + cond
q, k, v, o = self.self_attn(nx, nx, nx, emb_fn=None,
attn_mask=None,
key_padding_mask=None,
need_weights=False, is_causal=False, return_qkv=True)
x = x + self.layer_scale_1(self.dropout1(o))
x = x + self.layer_scale_2(self._ff_block(self.norm2(x)))
return q, k, v, x
class CPTransformer(nn.Module):
def __init__(self, model, emb_fn, start_layer, latent_dim, autocast, stride=50 * 10):
super().__init__()
self.emb_fn = {
"emb": emb_fn
}
new_layers = nn.ModuleList()
hidden_dim = 2048
cond_dim = latent_dim
num_layers = len(model.layers) - start_layer
max_n_frames = 500
self.pos_emb = nn.Parameter(
torch.randn(num_layers + 1, max_n_frames + 1, hidden_dim),
requires_grad=True)
# self.encodec_emb = nn.Linear(hidden_dim, latent_dim, bias=False)
self.merge_linear = nn.ModuleList()
# self.piano_roll_emb = nn.ModuleList()
for i in range(start_layer, len(model.layers)):
norm1 = model.layers[i].norm1
norm2 = model.layers[i].norm2
layer_scale_1 = model.layers[i].layer_scale_1
dropout1 = model.layers[i].dropout1
self_attn = model.layers[i].self_attn
layer_scale_2 = model.layers[i].layer_scale_2
linear1 = model.layers[i].linear1
linear2 = model.layers[i].linear2
activation = model.layers[i].activation
dropout = model.layers[i].dropout
new_layers.append(CPTransformerLayer(norm1=norm1,
norm2=norm2,
layer_scale_1=layer_scale_1,
dropout1=dropout1,
self_attn=self_attn,
linear1=linear1,
linear2=linear2,
activation=activation,
dropout=dropout,
layer_scale_2=layer_scale_2,
autocast=autocast))
self.merge_linear.append(nn.Linear(cond_dim, hidden_dim, bias=False))
# self.piano_roll_emb.append(nn.Linear(128, latent_dim, bias=False))
self.layers = new_layers
# self.gates = nn.Parameter(torch.zeros([num_layers, max_n_frames, 64]))
self.gates = nn.Parameter(torch.zeros([num_layers]))
freeze(self.layers)
self.max_n_frames = max_n_frames
self.start_layer = start_layer
self.num_layers = num_layers
self.stride = stride
def fn(self, i, activates):
if i >= self.num_layers:
return None, None
return activates[i]
def forward(self, condition_audio_code, max_n_frames, mode, skip=None):
max_n_frames = self.max_n_frames if max_n_frames is None else max_n_frames
sum_code = sum([self.emb_fn["emb"][i](condition_audio_code[:, i]) for i in range(4)])
condition_audio_code = sum_code
# condition_audio_code = self.encodec_emb(sum_code)
B, T, latent_dim = condition_audio_code.shape # (batch_size, n_frames, latent_dim)
o = self.pos_emb[0][None, :T].repeat(B, 1, 1)
#print(o.shape, T)
outs = []
encoded_condition = condition_audio_code
for i in range(len(self.layers)):
# We conduct two pass transformer.
# The first pass is to get multi-layer representation of the condition_audio_code.
# The second pass is to fuse the condition back.
# 1st pass
# encoded_condition = encoded_condition + self.pos_emb[i + 1][None, :T].repeat(B, 1, 1)
# _, _, _, encoded_condition = self.layers[i](x=encoded_condition, cond=None)
# add positional encoding and send to the transformer
embedding = self.merge_linear[i](encoded_condition) + self.pos_emb[i + 1][None, :T].repeat(B, 1, 1)
# embedding = (encoded_condition) + self.pos_emb[i + 1][None, :T].repeat(B, 1, 1) + self.merge_linear[i](condition_audio_code)
# 2nd pass
q, k, v, o = self.layers[i](x=o, cond=embedding) #
if not mode == "train":
outs.append([[torch.cat([q, q], 0),
torch.cat([k, k], 0),
torch.cat([v, v], 0)], self.gates[i]])
else:
outs.append([[q, k, v], self.gates[i]])
emb_fn = EmbFn(activates=outs, fn=self.fn,
start_layer=self.start_layer,
max_len=max_n_frames,
inference=(mode == "inference"),
skip=skip)
return emb_fn
def save_weights(self, path):
state_dict = {}
sdict = self.state_dict()
for n in sdict:
if str.startswith(n, "layers"):
continue
state_dict[n] = sdict[n]
torch.save(state_dict, path)
def load_weights(self, path):
self.load_state_dict(torch.load(path, map_location="cpu"), strict=False)
class Instructor(nn.Module):
def __init__(self, sec, num_layers, latent_dim, top_k):
'''The MusicGen model with instructor adapter.
Args:
sec: int, duration of the audio in seconds
num_layers: int, number of layers of adapter in the transformer
latent_dim: int, dimension of the latent space
'''
super().__init__()
lm = CondMusicgen(sec, top_k=top_k)
self.peft_model = lm
self.musicgen = lm.musicgen
self.cp_transformer = CPTransformer(self.musicgen.lm.transformer,
emb_fn=self.musicgen.lm.emb,
start_layer=48 - num_layers,
latent_dim=latent_dim,
autocast=self.musicgen.autocast)
self.text_lora_config = peft.LoraConfig(target_modules=r".*\.cross_attention\.(q_proj|v_proj)",
r=32,
lora_alpha=64)
self.peft_model.lm.transformer = peft.get_peft_model(self.peft_model.lm.transformer, self.text_lora_config)
def set_training(self):
self.peft_model.set_training()
print_trainable_parameters(self)
def save_weights(self, path):
self.cp_transformer.save_weights(path)
def load_weights(self, path):
self.cp_transformer.load_weights(path)
def forward(self, input_code, text_description, condition_audio_code,
num_samples=8, mode="train", max_n_frames=None, prompt_tokens=None):
if max_n_frames is None:
max_n_frames = input_code.shape[-1]
condition_audio_code = torch.cat([condition_audio_code, torch.ones_like(condition_audio_code[:,:, 0:1]) * 2048], dim=-1)
embed_fn = self.cp_transformer.forward(condition_audio_code=condition_audio_code,
max_n_frames=max_n_frames,
mode=mode,
skip=None)
out = self.peft_model.forward(input_code,
text_description=text_description,
embed_fn=embed_fn,
mode=mode,
total_gen_len=max_n_frames,
prompt_tokens=prompt_tokens)
return out
def generate(self, text_description, condition_audio_code,
num_samples=1):
out = self.peft_model.generate(cp_fn=self.cp_transformer,
text_description=text_description,
condition_audio_code=condition_audio_code,
num_samples=num_samples,
)
return out