File size: 7,698 Bytes
d911efa | 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 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 | """Bounded random access to indexed JSONL metadata and existing codec shards.
Only offsets are memory-mapped. No corpus-sized Python list or concatenated
audio/code tensor is copied into each GPU process. TAR member indices and
Parquet reference tables are kept in small per-thread LRU caches.
"""
from collections import OrderedDict
import hashlib
import io
import json
import os
from pathlib import Path
import tarfile
import threading
import numpy as np
class ShardReader:
def __init__(self, capacity=16, file_capacity=None):
import resource
self.capacity = capacity
soft_limit = resource.getrlimit(resource.RLIMIT_NOFILE)[0]
self.file_capacity = file_capacity or max(32, min(2048, soft_limit // 2))
self.local = threading.local()
def cache(self):
if not hasattr(self.local, 'files'):
self.local.files = OrderedDict()
self.local.tables = OrderedDict()
self.local.members = OrderedDict()
return self.local
def bytes_at(self, path, offset, size):
state = self.cache()
if path not in state.files:
state.files[path] = os.open(path, os.O_RDONLY)
while len(state.files) > self.file_capacity:
_, fd = state.files.popitem(last=False)
os.close(fd)
fd = state.files[path]
state.files.move_to_end(path)
value = os.pread(fd, int(size), int(offset))
if len(value) != size:
raise IOError(f'Short shard read: {path}@{offset}+{size}')
return value
def member(self, path, name):
state = self.cache()
if path not in state.members:
with tarfile.open(path, 'r:') as archive:
state.members[path] = {m.name: (m.offset_data, m.size) for m in archive if m.isfile()}
while len(state.members) > self.capacity:
state.members.popitem(last=False)
state.members.move_to_end(path)
offset, size = state.members[path][name]
return self.bytes_at(path, offset, size)
def codes(self, locator):
kind = locator.get('kind', 'tar_moss_npy')
if kind == 'tar_moss_npy':
blob = (self.bytes_at(locator['path'], locator['offset'], locator['size'])
if 'offset' in locator else self.member(locator['path'], locator['member']))
try:
result = np.load(io.BytesIO(blob), allow_pickle=False)
except ValueError:
# Raw little-endian uint16 dump without NumPy header
# (sidecar packer wrote tobytes()); shape from byte count.
result = np.frombuffer(blob, dtype=np.uint16).reshape(-1, 12)
assert result.dtype == np.uint16
elif kind == 'local_moss_parquet':
import pyarrow.parquet as pq
state = self.cache()
key = (locator['path'], locator['codes_field'], locator['frames_field'])
if key not in state.tables:
state.tables[key] = pq.read_table(locator['path'], columns=[locator['codes_field'], locator['frames_field']])
while len(state.tables) > self.capacity:
state.tables.popitem(last=False)
state.tables.move_to_end(key)
row = state.tables[key].slice(locator['row_index'], 1).to_pylist()[0]
frames = int(row[locator['frames_field']])
result = np.frombuffer(row[locator['codes_field']], dtype='<i2').reshape(frames, 12)
else:
raise ValueError(f'Unsupported verified codec locator: {kind}')
assert result.ndim == 2 and result.shape[1] == 12 and len(result) > 0
assert int(result.min()) >= 0 and int(result.max()) < 1024
expected = locator.get('full_frames')
if expected is not None:
assert len(result) == expected
if 'crop_frames' in locator:
start, length = int(locator.get('crop_start_frame', 0)), int(locator['crop_frames'])
assert start >= 0 and length > 0 and start + length <= len(result)
result = result[start:start + length]
return result
def record(self, locator):
blob = (self.bytes_at(locator['path'], locator['offset'], locator['size'])
if 'offset' in locator else self.member(locator['path'], locator['member']))
return json.loads(blob)
class IndexedRecords:
def __init__(self, manifest):
self.manifest_path = Path(manifest)
self.spec = json.loads(self.manifest_path.read_text())
root = self.manifest_path.parent
self.offsets = np.load(root / self.spec['offsets'], mmap_mode='r', allow_pickle=False)
self.path = root / self.spec['records']
self.fd = os.open(self.path, os.O_RDONLY)
assert len(self.offsets) == self.spec['samples'] + 1
assert int(self.offsets[0]) == 0 and int(self.offsets[-1]) == self.path.stat().st_size
self.reader = ShardReader()
def __len__(self):
return len(self.offsets) - 1
def row(self, index):
start, end = int(self.offsets[index]), int(self.offsets[index + 1])
return json.loads(os.pread(self.fd, end - start, start))
def example(self, index, packer):
from modes import verify_record_mode
row = self.row(index)
verify_record_mode(row)
mode = row['mode']
meta = self.reader.record(row['record_locator'])
assert meta['uid'] == row['uid']
meta['frames'] = int(meta['moss_frames'])
codes = self.reader.codes(row['target_locator'])
assert len(codes) == meta['frames']
reference = None
if mode == 'reference':
assert row['reference_uid'] != row['uid'] and row['reference_evidence']
reference = self.reader.codes(row['reference_locator'])
assert 0 < len(reference) <= 37
elif mode == 'score':
meta['conditioning_measurements'] = row['conditioning_measurements']
result = packer.pack_mode(meta, codes, mode, reference)
result['accounting'] = {'uid': row['uid'], 'mode': mode, 'frames': len(codes),
'target_audio_hours': float(meta['dur_s']) / 3600,
'reference_frames': len(reference) if reference is not None else 0}
return result
def write_manifest(directory, rows, contract):
directory = Path(directory)
directory.mkdir(parents=True, exist_ok=True)
offsets, n, modes, duration = [0], 0, {}, 0.
path = directory / 'records.jsonl'
digest = hashlib.sha256()
seen = set()
with path.with_suffix('.jsonl.tmp').open('wb') as stream:
for row in rows:
assert row['uid'] not in seen, row['uid']
seen.add(row['uid'])
assert row['mode'] in ('instruction', 'reference', 'score')
blob = (json.dumps(row, ensure_ascii=False, separators=(',', ':')) + '\n').encode()
stream.write(blob); digest.update(blob); offsets.append(offsets[-1] + len(blob)); n += 1
modes[row['mode']] = modes.get(row['mode'], 0) + 1
duration += row.get('dur_s', 0.) / 3600
path.with_suffix('.jsonl.tmp').replace(path)
np.save(directory / 'offsets.npy', np.asarray(offsets, dtype=np.int64), allow_pickle=False)
result = dict(contract, records=path.name, offsets='offsets.npy', samples=n, modes=modes,
target_audio_hours=duration, records_sha256=digest.hexdigest())
temporary = directory / 'manifest.json.tmp'
temporary.write_text(json.dumps(result, indent=2) + '\n')
temporary.replace(directory / 'manifest.json')
return directory / 'manifest.json'
|