Spaces:
Sleeping
Sleeping
File size: 4,050 Bytes
fe7e262 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 | import numpy as np
import torch
from tqdm import tqdm
from .model import PiCoGenDecoder
from .repr import Event
from .utils import downbeat_time_to_index
@torch.no_grad()
def decode(
model,
tokenizer,
beat_information,
melody_last_embs,
harmony_last_embs,
max_bar_num=None,
max_token_num=None,
temperature=1.0,
device=None,
):
if device is None:
device = model.parameters().__next__().device
starting_bpm = 60 / np.diff(np.array(beat_information["beats"])[:2]).mean()
TARGET = PiCoGenDecoder.InputClass.TARGET.value
CONDITION = PiCoGenDecoder.InputClass.CONDITION.value
out_events = [
Event(etype="spec", value="spec_ss"),
tokenizer.get_tempo_event(starting_bpm),
]
input_seg = [tokenizer.e2i(Event(etype="spec", value="spec_bos"))]
input_seg.extend([tokenizer.e2i(e) for e in out_events]) # song start, global tempo
need_encode_seg = [0] * len(input_seg)
input_cls_seg = [TARGET] * len(input_seg)
end_event = Event(etype="spec", value="spec_se")
bar_start_event = Event(etype="bar", value="bar_start")
bar_end_event = Event(etype="bar", value="bar_end")
total_beats = len(beat_information["beats"])
downbeats = downbeat_time_to_index(beat_information["beats"], beat_information["downbeats"])
if downbeats[-1] < total_beats:
downbeats.append(total_beats - 1)
if max_bar_num is not None:
downbeats = downbeats[: max_bar_num + 1]
pbar = tqdm(total=len(downbeats) - 1)
for bar_i, b in enumerate(range(len(downbeats) - 1)):
last_past_kv = None
# NOTE: upbeat is handled by SheetSage
downbeat_start, downbeat_end = downbeats[b], downbeats[b + 1]
# downbeat_start, downbeat_end = downbeats[b]-downbeats[0], downbeats[b+1]-downbeats[0]
for j in range(downbeat_start * tokenizer.beat_div, downbeat_end * tokenizer.beat_div):
input_seg.append((melody_last_embs[j], harmony_last_embs[j]))
input_cls_seg.append(CONDITION)
need_encode_seg.append(1)
if b == len(downbeats) - 2: # NOTE: add song_end to the last bar condition
input_seg.append(tokenizer.e2i(Event(etype="spec", value="spec_se")))
input_cls_seg.append(CONDITION)
need_encode_seg.append(0)
input_seg.append(tokenizer.e2i(bar_start_event))
need_encode_seg.append(0)
input_cls_seg.append(TARGET)
out_events.append(bar_start_event)
while True: # generate one bar
if len(input_seg) > model.hp.max_seq_len:
input_seg = input_seg[-model.hp.max_seq_len // 2 :]
input_seg[0] = tokenizer.e2i(Event(etype="spec", value="spec_bos"))
need_encode_seg = need_encode_seg[-model.hp.max_seq_len // 2 :]
need_encode_seg[0] = 0
input_cls_seg = input_cls_seg[-model.hp.max_seq_len // 2 :]
input_cls_seg[0] = TARGET
last_past_kv = None
input_cls_ids = torch.LongTensor(input_cls_seg)[None, :].to(device)
need_encode = torch.BoolTensor(need_encode_seg)[None, :].to(device)
output_ids, past_kv = model.generate(
input_seg=[input_seg],
input_cls_ids=input_cls_ids,
need_encode=need_encode,
kv_cache=last_past_kv,
temperature=temperature,
)
out_id = output_ids[0][-1].item()
out_event = tokenizer.i2e(out_id)
out_events.append(out_event)
input_seg.append(out_id)
input_cls_seg.append(TARGET)
need_encode_seg.append(0)
if out_event in (bar_end_event, end_event):
break
last_past_kv = past_kv
pbar.set_description(f"length: {len(out_events)}({len(input_seg)})")
pbar.update(1)
if max_token_num is not None and len(input_seg) > max_token_num:
break
pbar.close()
return out_events
|