File size: 5,413 Bytes
cd9b2d8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Frozen selections and targets from existing Ladder/P3/Gemini readers."""
import hashlib
import math
import sys
import numpy as np
from study_paths import CODE, ROOT, GEMINI, RELEASE, SEED, read

sys.path.insert(0, str(CODE))
from train_full_curriculum import normalization, TIMBRE_INDEX, IDENTITY_INDEX
from layered_curriculum_data import LadderAuxDataset, taxonomy
from p3_frozen_loader import P3FrozenDataset
from timbre_targets import TimbreIndex


class SourceData:
    def __init__(self, domain):
        self.domain = domain
        self.mean, self.std, self.names = normalization()
        self.norm = read(RELEASE / 'training_normalization.json')
        self.classes = read(RELEASE / 'classes.json')
        self.map = self.classes if 'ladder_raw_to_canonical' in self.classes else taxonomy()
        if domain == 'ladder':
            self.data = LadderAuxDataset(self.mean, self.std)
            self.indices = np.asarray([i for s, shard in enumerate(self.data.shards)
                                       if shard['level'] in ('t0', 't1', 't2')
                                       for i in range(0 if s == 0 else int(self.data.ends[s-1]), int(self.data.ends[s]))], np.int64)
            self.speakers = {'timbre': TimbreIndex(TIMBRE_INDEX, 128), 'identity': TimbreIndex(IDENTITY_INDEX, 250)}
        elif domain.startswith('p3_'):
            split = domain[3:]
            self.data = P3FrozenDataset(split, score_mean=self.mean, score_std=self.std,
                                       event_class_path=__import__('pathlib').Path('/e/scratch/reformo/schuhmann1_moss/vocalburst_p3/frozen_1120k/event_class_ids.npy'))
            # Enough training clips for all three 10% stage mixes. Validation
            # and test use the same fixed 2k-clips-per-split pilot across models.
            size = min(len(self.data), 22000 if split == 'train' else 2000)
            self.indices = np.random.default_rng(SEED).choice(len(self.data), size, replace=False)
        elif domain == 'gemini':
            sys.path.insert(0, str(CODE / 'gemini_finetune'))
            from common import TarReader
            self.reader = TarReader()
            selection = GEMINI.parent / 'gemini_multisource_100h_20261004/combined_unique_manifest.jsonl'
            self.data = [__import__('json').loads(line) for line in selection.open()]
            self.indices = np.arange(len(self.data))
        else:
            from evaluate_public_benchmarks import parquet_samples, emolia_samples
            self.data = parquet_samples(domain) if domain in ('emonet', 'crema', 'ravdess') else emolia_samples(domain)
            self.indices = np.arange(len(self.data))

    def __len__(self):
        return len(self.indices)

    def __getitem__(self, position):
        index = int(self.indices[position])
        if self.domain == 'gemini':
            row = self.data[index]
            return row['sha256_mp3'], self.reader.audio(row, verify=True), {'source': 'gemini', 'key': row['sha256_mp3']}, {}
        if self.domain not in ('ladder', 'p3_train', 'p3_validation', 'p3_test'):
            from evaluate_public_benchmarks import decode
            sample = self.data[index]
            wave, duration, truncated = decode(sample)
            return sample.key, wave, {**sample.metadata, 'source': self.domain, 'key': sample.key,
                                     'duration_s': duration, 'truncated_30s': truncated}, {}
        item = self.data[index]
        if self.domain == 'ladder':
            shard = int(np.searchsorted(self.data.ends, index, side='right'))
            level = self.data.shards[shard]['level']
            stages = ['S8'] + (['S9'] if level in ('t0', 't1') else []) + (['S10'] if level == 't0' else [])
            value = int(hashlib.sha256(item['uid'].encode()).hexdigest()[:8], 16) % 100
            split = 'train' if value < 90 else 'validation' if value < 95 else 'test'
            lookup = self.map['ladder_raw_to_canonical']
            spans = item['events_frames']
            for name, speaker in self.speakers.items():
                vector, valid = speaker.lookup([item['uid']])
                item[name], item[name + '_valid'] = vector[0].astype(np.float32), bool(valid[0])
        else:
            stages = ['S8', 'S9', 'S10']
            split = self.domain[3:]
            lookup = self.map['p3_raw_to_canonical']
            spans = [(start // 320, math.ceil(end / 320)) for start, end in item['events_samples']]
            for name in ('timbre', 'identity'):
                item[name + '_valid'] = item['speaker_valid']
        classes = [lookup[int(c)] for c in item['event_raw_class_ids'][:len(spans)]]
        target = {k: np.asarray(item[k]) for k in ('scores', 'score_mask', 'frame', 'frame_mask', 'timbre', 'identity', 'timbre_valid', 'identity_valid')}
        cps = item.get('cps_raw')
        cps_valid = cps is not None and np.isfinite(cps)
        target.update(cps=np.float32((cps - self.norm['cps_mean']) / self.norm['cps_std']) if cps_valid else np.float32(0),
                      cps_valid=np.bool_(cps_valid), event_starts=np.asarray([x[0] for x in spans], np.int64),
                      event_ends=np.asarray([x[1] for x in spans], np.int64), event_classes=np.asarray(classes, np.int64))
        return self.domain + ':' + item['key'], item['audio'], {'source': self.domain, 'key': item['key'], 'split': split, 'stages': stages}, target