Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions ariautils/config/config.json
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,7 @@
"time_step_ms": 10,
"include_drums": true,
"include_pedal": false,
"include_delimiter": false,
"composer_names": ["bach", "beethoven", "mozart", "chopin", "rachmaninoff", "liszt", "debussy", "schubert", "brahms", "ravel", "satie", "scarlatti"],
"form_names": ["sonata", "prelude", "nocturne", "étude", "waltz", "mazurka", "impromptu", "fugue"],
"genre_names": ["jazz", "classical"]
Expand Down
31 changes: 29 additions & 2 deletions ariautils/tokenizer/absolute.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,7 +100,7 @@ def __init__(self, config_path: Path | str | None = None) -> None:
if v is False
]

self.include_drums: bool = self.config.get("include_drums", True)
self.include_drums = self.config.get["include_drums"]
if self.include_drums:
self.instruments_wd = self.instruments_nd + ["drum"]
else:
Expand Down Expand Up @@ -156,9 +156,15 @@ def __init__(self, config_path: Path | str | None = None) -> None:
self.include_pedal = self.config["include_pedal"]
self.ped_on_tok = "<PED_ON>"
self.ped_off_tok = "<PED_OFF>"
if self.config["include_pedal"] is True:
if self.include_pedal is True:
self.add_tokens_to_vocab([self.ped_on_tok, self.ped_off_tok])

self.include_delimiter = self.config["include_delimiter"]
self.delimiter_tok = "<X>"
if self.include_delimiter is True:
self.add_tokens_to_vocab([self.delimiter_tok])
self.special_tokens.append(self.delimiter_tok)

def export_data_aug(self) -> list[Callable[[list[Token]], list[Token]]]:
return [
self.export_tempo_aug(max_tempo_aug=0.2, mixup=True),
Expand Down Expand Up @@ -878,6 +884,7 @@ def export_tempo_aug(
Callable[[list[Token], float], list[Token]]: Exported function.
"""

# TODO: Potential issue with delimiter_tok at start
def tempo_aug(
src: list[Token],
abs_time_step: int,
Expand All @@ -887,6 +894,7 @@ def tempo_aug(
eos_tok: str,
time_tok: str,
dim_tok: str,
delimiter_tok: str,
pad_tok: str,
unk_tok: str,
ped_on_tok: str,
Expand Down Expand Up @@ -917,6 +925,7 @@ def _quantize_time(_n: int | float) -> int:
res_prefix: list[Token] = []
src_time_tok_cnt = 0
dim_tok_seen_at: tuple[int, int] | None = None
delimiter_tok_seen_at: tuple[int, int] | None = None
eos_tok_seen: bool = False

idx = 0
Expand Down Expand Up @@ -962,6 +971,20 @@ def _quantize_time(_n: int | float) -> int:
dim_tok_seen_at = (last_time, last_onset)
idx += 1
continue
if tok == delimiter_tok:
if delimiter_tok_seen_at is not None:
logger.warning(
"Multiple <X> tokens encountered in augmentation"
)
last_time = max(buffer.keys()) if buffer else 0
last_onset = (
max(buffer[last_time].keys())
if buffer.get(last_time)
else 0
)
delimiter_tok_seen_at = (last_time, last_onset)
idx += 1
continue

# Parse event sequences (note, drum, pedal)
tok_type = tok[0] if isinstance(tok, tuple) else tok
Expand Down Expand Up @@ -1069,6 +1092,9 @@ def _quantize_time(_n: int | float) -> int:
if dim_tok_seen_at == (src_time_tok_cnt, src_onset):
res_events.append(dim_tok)
dim_tok_seen_at = None
if delimiter_tok_seen_at == (src_time_tok_cnt, src_onset):
res_events.append(delimiter_tok)
delimiter_tok_seen_at = None

# Re-assemble the final sequence
final_res = res_prefix + res_events
Expand All @@ -1088,6 +1114,7 @@ def _quantize_time(_n: int | float) -> int:
eos_tok=self.eos_tok,
time_tok=self.time_tok,
dim_tok=self.dim_tok,
delimiter_tok=self.delimiter_tok,
pad_tok=self.pad_tok,
unk_tok=self.unk_tok,
ped_on_tok=self.ped_on_tok,
Expand Down
Loading