Views
No views yet
[0.123, 5.546]) are extracted from data and model outputs
1import numpy as np
2
3from datasets import Audio, Dataset
4from transformers import AutoFeatureExtractor, Wav2Vec2BertForAudioFrameClassification
5import torch
6import numpy as np
7
8if torch.cuda.is_available():
9 device = torch.device("cuda")
10else:
11 device = torch.device("cpu")
12
13model_name = "classla/wav2vecbert2-prosodicUnit"
14feature_extractor = AutoFeatureExtractor.from_pretrained(model_name)
15model = Wav2Vec2BertForAudioFrameClassification.from_pretrained(model_name).to(device)
16f = "data/Rog-Art-N-G6007-P600702_181.070_211.070.wav"
17
18
19def frames_to_intervals(frames: list) -> list[tuple]:
20 from itertools import pairwise
21 import pandas as pd
22
23 results = []
24 ndf = pd.DataFrame(
25 data={
26 "time_s": [0.020 * i for i in range(len(frames))],
27 "frames": frames,
28 }
29 )
30 ndf = ndf.dropna()
31 indices_of_change = ndf.frames.diff()[ndf.frames.diff() != 0].index.values
32 for si, ei in pairwise(indices_of_change):
33 if ndf.loc[si : ei - 1, "frames"].mode()[0] == 0:
34 pass
35 else:
36 results.append(
37 (round(ndf.loc[si, "time_s"], 3), round(ndf.loc[ei - 1, "time_s"], 3))
38 )
39 return results
40
41
42def evaluator(chunks):
43 sampling_rate = chunks["audio"][0]["sampling_rate"]
44 with torch.no_grad():
45 inputs = feature_extractor(
46 [i["array"] for i in chunks["audio"]],
47 return_tensors="pt",
48 sampling_rate=sampling_rate,
49 ).to(device)
50 logits = model(**inputs).logits
51 y_pred_raw = np.array(logits.cpu())
52 y_pred = y_pred_raw.argmax(axis=-1)
53 prosodic_units = [frames_to_intervals(i) for i in y_pred]
54 return {
55 "y_pred": y_pred,
56 "y_pred_logits": y_pred_raw,
57 "prosodic_units": prosodic_units,
58 }
59
60# Create a dataset with a single instance and map our evaluator function on it:
61ds = Dataset.from_dict({"audio": [f]}).cast_column("audio", Audio(16000, mono=True))
62ds = ds.map(evaluator, batched=True, batch_size=1) # Adjust batch size according to your hardware specs
63print(ds["y_pred"][0])
64# Outputs: [0, 0, 1, 1, 1, 1, 1, ...]
65print(ds["y_pred_logits"][0])
66# Outputs:
67# [[ 0.89419061, -0.77746612],
68# [ 0.44213724, -0.34862748],
69# [-0.08605709, 0.13012762],
70# ....
71print(ds["prosodic_units"][0])
72# Outputs: [[0.04, 2.4], [3.52, 6.6], ....1import numpy as np
2
3from datasets import Audio, Dataset
4from transformers import AutoFeatureExtractor, Wav2Vec2BertForAudioFrameClassification
5import torch
6import numpy as np
7
8if torch.cuda.is_available():
9 device = torch.device("cuda")
10else:
11 device = torch.device("cpu")
12
13model_name = "classla/wav2vecbert2-prosodicUnit"
14feature_extractor = AutoFeatureExtractor.from_pretrained(model_name)
15model = Wav2Vec2BertForAudioFrameClassification.from_pretrained(model_name).to(device)
16f = "ROG/ROG-Art/WAV/Rog-Art-N-G5025-P600022.wav"
17
18OVERLAP_S = 10
19CHUNK_LENGTH_S = 30
20SAMPLING_RATE = 16_000
21OVERLAP_SAMPLES = OVERLAP_S * SAMPLING_RATE
22CHUNK_LENGTH_SAMPLES = CHUNK_LENGTH_S * SAMPLING_RATE
23
24
25def frames_to_intervals(frames: list) -> list[tuple]:
26 from itertools import pairwise
27 import pandas as pd
28
29 results = []
30 ndf = pd.DataFrame(
31 data={
32 "time_s": [0.020 * i for i in range(len(frames))],
33 "frames": frames,
34 }
35 )
36 ndf = ndf.dropna()
37 indices_of_change = ndf.frames.diff()[ndf.frames.diff() != 0].index.values
38 for si, ei in pairwise(indices_of_change):
39 if ndf.loc[si : ei - 1, "frames"].mode()[0] == 0:
40 pass
41 else:
42 results.append(
43 (round(ndf.loc[si, "time_s"], 3), round(ndf.loc[ei - 1, "time_s"], 3))
44 )
45 return results
46
47
48def merge_events(events: list[list[float]], centroids):
49 flattened_events = []
50 flattened_centroids = []
51 for batch_idx, batch in enumerate(events):
52 for event in batch:
53 flattened_events.append(event)
54 flattened_centroids.append(centroids[batch_idx])
55 flattened_events.sort(key=lambda x: x[0])
56
57 # Merged list to store final intervals
58 merged = []
59
60 for event, centroid in zip(flattened_events, flattened_centroids):
61 if not merged:
62 # If merged is empty, simply add the first event
63 merged.append((event, centroid))
64 else:
65 last_event, last_centroid = merged[-1]
66 # Check for overlap
67 if (last_event[0] < event[1]) and (last_event[1] > event[0]):
68 # Calculate the midpoint of the intervals
69 last_event_midpoint = (last_event[0] + last_event[1]) / 2
70 current_event_midpoint = (event[0] + event[1]) / 2
71
72 # Choose the event whose centroid is closer to its midpoint
73 if abs(last_centroid - last_event_midpoint) <= abs(
74 centroid - current_event_midpoint
75 ):
76 continue
77 else:
78 merged[-1] = (event, centroid)
79 else:
80 merged.append((event, centroid))
81
82 final_intervals = [event for event, _ in merged]
83 return final_intervals
84
85
86def evaluator(chunks):
87 with torch.no_grad():
88 samples = []
89 for array, start, end in zip(chunks["audio"], chunks["start"], chunks["end"]):
90 samples.append(array["array"][start:end])
91 inputs = feature_extractor(
92 samples,
93 return_tensors="pt",
94 sampling_rate=SAMPLING_RATE,
95 ).to(device)
96 logits = model(**inputs).logits
97 y_pred_raw = np.array(logits.cpu())
98 y_pred = y_pred_raw.argmax(axis=-1)
99 prosodic_units = [
100 np.array(frames_to_intervals(i)) + start / SAMPLING_RATE
101 for i, start in zip(y_pred, chunks["start"])
102 ]
103 return {
104 "y_pred": y_pred,
105 "y_pred_logits": y_pred_raw,
106 "prosodic_units": prosodic_units,
107 }
108
109
110audio_duration_samples = (
111 Audio(SAMPLING_RATE, mono=True)
112 .decode_example({"path": f, "bytes": None})["array"]
113 .shape[0]
114)
115chunk_starts = np.arange(
116 0, audio_duration_samples, CHUNK_LENGTH_SAMPLES - OVERLAP_SAMPLES
117)
118chunk_ends = chunk_starts + CHUNK_LENGTH_SAMPLES
119
120ds = Dataset.from_dict(
121 {
122 "audio": [f for i in chunk_starts],
123 "start": chunk_starts,
124 "end": chunk_ends,
125 "chunk_centroid_s": (chunk_starts + chunk_ends) / 2 / SAMPLING_RATE,
126 }
127).cast_column("audio", Audio(SAMPLING_RATE, mono=True))
128
129ds = ds.map(evaluator, batched=True, batch_size=10)
130
131
132final_intervals = merge_events(ds["prosodic_units"], ds["chunk_centroid_s"])
133print(final_intervals)
134# Outputs: [[3.14, 4.96], [5.6, 8.4], [8.62, 9.32], [10.12, 10.7], [11.72, 13.1],....| hyperparameter | value |
|---|---|
| learning rate | 3e-5 |
| effective batch size | 16 |
| num train epochs | 20 |
mamba create -f transformers_env.yml (replace mamba with conda if you don't
use mamba).