File size: 7,933 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
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
"""Read cached native features and attach the selected target regime."""
import io
import json
import math
import os
from collections import OrderedDict
from pathlib import Path
import numpy as np
import torch
from study_paths import ROOT, GEMINI, RELEASE, read

_COMMITTED_INDEX_ROWS = {}


class FeatureData:
    def __init__(self, model, domains=None, split=None, phase='legacy', stage=None):
        self.rows = []
        self.fds = OrderedDict()
        self.projection = None
        self.phase = phase
        self.flash = None
        self.norm = read(RELEASE / 'training_normalization.json')
        if phase == 'gemini':
            if not (GEMINI / 'prepared/TRAINING_READY.json').exists():
                raise RuntimeError('Final Gemini training gate is not ready')
            self.flash = {r['sha256_mp3']: r for r in map(json.loads, (GEMINI / 'prepared/targets.jsonl').open())
                          if r['ready_for_whisper_training'] and (split is None or r['split'] == split)}
        feature_root=ROOT/'features'/model
        domain_key=tuple(sorted(domains)) if domains else ()
        cache_key=(model,domain_key)
        committed=bool(domains) and all((feature_root/(d+'-rank'+str(r)+'-COMPLETE.json')).exists()
                                      for d in domains for r in range(4))
        # The same immutable index serves train/validation/test and S8/S9/S10.
        # Decode its metadata once per process; keep the original split filters.
        indexed=_COMMITTED_INDEX_ROWS.get(cache_key) if committed else None
        if indexed is None:
            indexed=[]
            for path in sorted(feature_root.glob('*.jsonl')):
                if domains and not any(path.name.startswith(d+'-rank') for d in domains):
                    continue
                with path.open() as stream:
                    indexed.extend(row for row in map(json.loads,stream)
                                   if not domains or row['source'] in domains)
            if committed:_COMMITTED_INDEX_ROWS[cache_key]=indexed
        for row in indexed:
            if phase == 'gemini':
                if row['source'] != 'gemini' or row['key'] not in self.flash:
                    continue
            elif split and row.get('split') != split:
                continue
            if stage and stage not in row.get('stages', []):
                continue
            self.rows.append(row)

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

    def payload(self, row):
        path = row['feature_tar']
        if path not in self.fds:
            if len(self.fds) >= 24:
                _, old = self.fds.popitem(last=False)
                os.close(old)
            self.fds[path] = os.open(path, os.O_RDONLY)
        self.fds.move_to_end(path)
        raw = os.pread(self.fds[path], row['feature_size'], row['feature_offset'])
        if len(raw) != row['feature_size']:
            raise RuntimeError('Short feature cache read')
        with np.load(io.BytesIO(raw), allow_pickle=False) as data:
            return {k: data[k] for k in data.files}

    def __getstate__(self):
        state = {**self.__dict__, 'fds': OrderedDict()}
        state.pop('teacher_reader', None)
        return state

    def __getitem__(self, index):
        row = self.rows[index]
        value = self.payload(row)
        embedding = value['embedding'].astype(np.float32)
        if self.projection:
            p = self.projection
            embedding = ((embedding - p['mean']) @ p['components'].T) / p['scale']
            embedding = np.pad(embedding, (0, 256 - len(embedding))).astype(np.float32)
        frames = int(math.ceil(float(value['duration_s']) * 50))
        # The model sees interpolated native features on the shared 20-ms grid.
        # Native timestamps/count are retained for resolution-aware reporting.
        ticks = (np.arange(frames) + .5) / 50
        native = value['frame_features'].astype(np.float32)
        times = value['frame_times_s']
        value['frame_features'] = np.stack([np.interp(ticks, times, native[:, i]) for i in range(64)], axis=1).astype(np.float32)
        value['embedding'] = embedding
        value['feature_key'] = row['feature_key']
        if self.phase == 'gemini':
            # Targets are prepared without decoding the cached full audio.
            r = self.flash[row['key']]
            norm = self.norm
            raw = np.asarray([r['raw_scores'].get(n, np.nan) for n in norm['score_names']], np.float32)
            valid = np.isfinite(raw)
            teachers = r.get('teacher_metadata', {})
            if not teachers.get('dnsmos', {}).get('valid_for_speech', False):
                valid[123:130] = False
            if not teachers.get('empathic', {}).get('speech_domain_valid', False):
                valid[100:119] = False
            if not teachers.get('voiceclap', {}).get('speech_domain_valid', False):
                for i in range(131, 192):
                    if r['score_provenance'].get(norm['score_names'][i]) != 'gemini-3.8-flash':
                        valid[i] = False
            value.update(scores=np.where(valid, (raw - norm['score_mean']) / norm['score_std'], 0).astype(np.float32),
                         score_mask=valid, frame=np.zeros(frames, np.float32),
                         frame_mask=np.full(frames, r['burst_frames_valid'], np.float32))
            events = r['vocal_bursts']
            starts = np.asarray([max(0, min(frames-1, int(e['start'] * 50))) for e in events], np.int64)
            ends = np.asarray([min(frames, max(int(a)+1, math.ceil(e['end'] * 50))) for a,e in zip(starts,events)], np.int64)
            for a,b in zip(starts,ends):
                value['frame'][a:b] = 1
            cps = r['characters_per_second']
            value.update(event_starts=starts, event_ends=ends,
                         event_classes=np.asarray([e['checkpoint_class_id'] for e in events], np.int64),
                         cps=np.float32((cps-norm['cps_mean'])/norm['cps_std']) if cps is not None else np.float32(0),
                         cps_valid=np.bool_(cps is not None))
            import sys
            from study_paths import CODE
            sys.path.insert(0, str(CODE / 'gemini_finetune'))
            from common import TarReader
            if not hasattr(self, 'teacher_reader'):
                self.teacher_reader = TarReader()
            for name,dim in [('timbre',128),('identity',250)]:
                reference = r['embeddings'].get(name)
                value[name] = np.zeros(dim,np.float32)
                if reference:
                    value[name] = np.load(io.BytesIO(self.teacher_reader.read(reference['tar'],reference['member'])),allow_pickle=False).astype(np.float32)
                value[name+'_valid'] = np.bool_(reference and teachers.get('orange',{}).get('valid_single_speaker_target'))
        return value


def collate(items):
    width = max(len(x['frame_features']) for x in items)
    events = max(1, max(len(x.get('event_starts', [])) for x in items))
    arrays = {}
    for name in ('embedding','scores','score_mask','timbre','identity','timbre_valid','identity_valid','cps','cps_valid'):
        arrays[name] = torch.from_numpy(np.stack([x[name] for x in items]))
    for name in ('frame','frame_mask','frame_features'):
        source = [np.pad(x[name], ((0,width-len(x[name])),(0,0)) if x[name].ndim == 2 else (0,width-len(x[name]))) for x in items]
        arrays[name] = torch.from_numpy(np.stack(source))
    for name in ('event_starts','event_ends','event_classes'):
        arrays[name] = torch.from_numpy(np.stack([np.pad(x[name], (0,events-len(x[name]))) for x in items]))
    mask = np.stack([np.arange(events)<len(x['event_starts']) for x in items])
    arrays['event_valid'] = torch.from_numpy(mask)
    arrays['event_span_valid'] = arrays['event_valid'].clone()
    arrays['keys'] = [x['feature_key'] for x in items]
    return arrays