-
Notifications
You must be signed in to change notification settings - Fork 6
Expand file tree
/
Copy pathwrapper.py
More file actions
62 lines (55 loc) · 1.97 KB
/
Copy pathwrapper.py
File metadata and controls
62 lines (55 loc) · 1.97 KB
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
from typing import Union, List, Tuple
import numpy as np
from ..speeches.speechset import SpeechSet
class IDWrapper(SpeechSet):
"""Speechset wrapper for auxiliary ids.
"""
def __init__(self, speechset: SpeechSet):
"""Initializer.
Args:
speechset: base speechset.
"""
super().__init__(speechset.reader)
# hold
self.speechset = speechset
def normalize(self,
ids: Union[int, List[int]],
text: str,
speech: np.ndarray) \
-> Tuple[Union[int, List[int]], Tuple[np.ndarray, np.ndarray]]:
"""Normalize datum with auxiliary ids.
Args:
ids: auxiliary ids.
text: transcription.
speech: [np.float32; [T]], speech in range (-1, 1).
Returns:
id and normalized datum.
"""
return ids, self.speechset.normalize(text, speech)
def collate(self,
bunch: List[Tuple[Union[int, List[int]],
Tuple[np.ndarray, np.ndarray]]]) \
-> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
"""Collate bunch of datum to the batch data.
Args:
bunch: B x [...], list of normalized inputs.
ids: auxiliary ids.
...: normalized datum.
Returns:
bunch data.
ids: [np.long; [B, ...]], auxiliary ids.
...: collated bunch.
"""
# [B, ...], auxiliary ids.
ids = self.collate_id([ids for ids, _ in bunch])
# collated bunch
return (ids, *self.speechset.collate([datum for _, datum in bunch]))
def collate_id(self, bunch: List[Union[int, List[int]]]) -> np.ndarray:
"""ID collator.
Args:
bunch: B x [...], list of ids.
Returns:
[np.long; [B, ...]], collated ids.
"""
# simple wrapping
return np.array(bunch, dtype=np.int64)