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