Uh oh!
There was an error while loading. Please reload this page.
forked from golbin/WaveNet
- Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathgenerate.py
More file actions
Latest commit
94 lines (62 loc) · 2.65 KB
/
Copy pathgenerate.py
File metadata and controls
94 lines (62 loc) · 2.65 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
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
"""
A script for WaveNet training
"""
importtorch
importlibrosa
importdatetime
importnumpyasnp
importwavenet.configasconfig
fromwavenet.modelimportWaveNet
importwavenet.utils.dataasutils
classGenerator:
def__init__(self, args):
self.args=args
self.wavenet=WaveNet(args.layer_size, args.stack_size,
args.in_channels, args.res_channels)
self.wavenet.load(args.model_dir, args.step)
@staticmethod
def_variable(data):
tensor=torch.from_numpy(data).float()
iftorch.cuda.is_available():
returntorch.autograd.Variable(tensor.cuda())
else:
returntorch.autograd.Variable(tensor)
def_make_seed(self, audio):
audio=np.pad([audio], [[0, 0], [self.wavenet.receptive_fields, 0], [0, 0]], 'constant')
ifself.args.sample_size:
seed=audio[:, :self.args.sample_size, :]
else:
seed=audio[:, :self.wavenet.receptive_fields*2, :]
returnseed
def_get_seed_from_audio(self, filepath):
audio=utils.load_audio(filepath, self.args.sample_rate)
audio_length=len(audio)
audio=utils.mu_law_encode(audio, self.args.in_channels)
audio=utils.one_hot_encode(audio, self.args.in_channels)
seed=self._make_seed(audio)
returnself._variable(seed), audio_length
def_save_to_audio_file(self, data):
data=data[0].cpu().data.numpy()
data=utils.one_hot_decode(data, axis=1)
audio=utils.mu_law_decode(data, self.args.in_channels)
librosa.output.write_wav(self.args.out, audio, self.args.sample_rate)
print('Saved wav file at {}'.format(self.args.out))
returnlibrosa.get_duration(y=audio, sr=self.args.sample_rate)
defgenerate(self):
outputs= []
inputs, audio_length=self._get_seed_from_audio(self.args.seed)
whileTrue:
new=self.wavenet.generate(inputs)
outputs=torch.cat((outputs, new), dim=1) iflen(outputs) elsenew
print('{0}/{1} samples are generated.'.format(len(outputs[0]), audio_length))
iflen(outputs[0]) >=audio_length:
break
inputs=torch.cat((inputs[:, :-len(new[0]), :], new), dim=1)
outputs=outputs[:, :audio_length, :]
returnself._save_to_audio_file(outputs)
if__name__=='__main__':
args=config.parse_args(is_training=False)
generator=Generator(args)
start_time=datetime.datetime.now()
duration=generator.generate()
print('Generate {0} seconds took {1}'.format(duration, datetime.datetime.now() -start_time))