diff --git a/nv_wavenet.cuh b/nv_wavenet.cuh index 5baacc0..1340894 100644 --- a/nv_wavenet.cuh +++ b/nv_wavenet.cuh @@ -285,6 +285,9 @@ class nvWavenetInfer { void setActivation(float* dst, float* src, size_t size) { gpuErrChk(cudaMemcpy(dst, src, size*sizeof(float), cudaMemcpyDefault)); } + void setActivation(half* dst, half* src, size_t size) { + gpuErrChk(cudaMemcpy(dst, src, size*sizeof(half), cudaMemcpyDefault)); + } void setActivation(half* dst, float* src, size_t size) { convert_float2half(dst, src, size); } @@ -419,6 +422,12 @@ class nvWavenetInfer { setActivation(m_Lh, Lh, m_maxSamples*m_numLayers*m_maxBatch*2*R); gpuErrChk(cudaMemcpy(m_outputSelectors, outputSelectors, m_maxSamples*m_maxBatch*sizeof(float), cudaMemcpyHostToDevice)); + } + void setInputs (half* Lh, float* outputSelectors) { + silenceInputs<<<1,256>>>(m_yInPrev, m_yInCur, m_maxBatch); + setActivation(m_Lh, Lh, m_maxSamples*m_numLayers*m_maxBatch*2*R); + gpuErrChk(cudaMemcpy(m_outputSelectors, outputSelectors, m_maxSamples*m_maxBatch*sizeof(float), cudaMemcpyHostToDevice)); + } void getXtOut(int layer, float* hXt) { getActivation(hXt, m_XtOut + layer*m_maxBatch*R, m_maxBatch*R); } diff --git a/pytorch/wavenet_infer.cu b/pytorch/wavenet_infer.cu index ea5716c..63e19e0 100644 --- a/pytorch/wavenet_infer.cu +++ b/pytorch/wavenet_infer.cu @@ -36,8 +36,11 @@ const int A = 256; const int R = 64; const int S = 256; typedef nvWavenetInfer MyWaveNet; +typedef nvWavenetInfer MyWaveNet_half; -std::shared_ptr make_wavenet(int sample_count, + +template +std::shared_ptr make_wavenet(int sample_count, int batch_size, float* embedding_prev, float* embedding_curr, @@ -55,7 +58,7 @@ std::shared_ptr make_wavenet(int sample_count, bool use_embed_tanh, int implementation ) { - std::shared_ptr wavenet(new MyWaveNet(num_layers, max_dilation, + std::shared_ptr wavenet(new WaveNetType(num_layers, max_dilation, batch_size, sample_count, implementation, use_embed_tanh)); @@ -84,14 +87,15 @@ std::shared_ptr make_wavenet(int sample_count, return wavenet; } -void infer(std::shared_ptr wavenet, - float* input_features, +template +void infer(std::shared_ptr wavenet, + void* input_features, int* samples, int sample_count, int batch_size) { Matrix outputSelectors(batch_size, sample_count); outputSelectors.randomize(0.5,1.0); - wavenet->setInputs(input_features, outputSelectors.data()); + wavenet->setInputs((T_data*)input_features, outputSelectors.data()); int batch_size_per_block = ((batch_size % 4) == 0) ? 4 : ((batch_size % 2) == 0) ? 2 : 1; assert(wavenet->run(sample_count, batch_size, samples, batch_size_per_block, true)); @@ -118,10 +122,12 @@ void wavenet_infer(int sample_count, float* conv_out_weight, float* conv_end_weight, int use_embed_tanh, - float* cond_input, + void* cond_input, bool cond_half, int implementation, int* samples) { - std::shared_ptr wavenet = make_wavenet(sample_count, + assert(samples); + if (cond_half) { + std::shared_ptr wavenet_half = make_wavenet(sample_count, batch_size, embedding_prev, embedding_curr, @@ -139,9 +145,30 @@ void wavenet_infer(int sample_count, use_embed_tanh, implementation ); - assert(samples); - infer(wavenet, cond_input, samples, sample_count, batch_size); - return; + infer(wavenet_half, cond_input, samples, sample_count, batch_size); + + } else { + std::shared_ptr wavenet = make_wavenet(sample_count, + batch_size, + embedding_prev, + embedding_curr, + num_layers, + max_dilation, + in_layer_weights_prev, + in_layer_weights_curr, + in_layer_biases, + res_layer_weights, + res_layer_biases, + skip_layer_weights, + skip_layer_biases, + conv_out_weight, + conv_end_weight, + use_embed_tanh, + implementation + ); + infer(wavenet, cond_input, samples, sample_count, batch_size); + } + } int get_R() {return R;} diff --git a/pytorch/wavenet_infer.h b/pytorch/wavenet_infer.h index a921749..cdb3d4a 100644 --- a/pytorch/wavenet_infer.h +++ b/pytorch/wavenet_infer.h @@ -30,6 +30,7 @@ extern "C" { // ------------------------------------------------ // C-compatible function for wrapper // ------------------------------------------------ + void wavenet_infer(int sample_count, int batch_size, float* embedding_prev, @@ -46,7 +47,7 @@ void wavenet_infer(int sample_count, float* conv_out_weight, float* conv_end_weight, int use_embed_tanh, - float* cond_input, + void* cond_input, bool half, int implementation, int* samples); diff --git a/pytorch/wavenet_infer_wrapper.cpp b/pytorch/wavenet_infer_wrapper.cpp index abafbcb..da1020f 100644 --- a/pytorch/wavenet_infer_wrapper.cpp +++ b/pytorch/wavenet_infer_wrapper.cpp @@ -48,7 +48,14 @@ int infer(at::Tensor samples_tensor, float* embedding_curr = embed_curr_tensor.data(); float* conv_out = conv_out_tensor.data(); float* conv_end = conv_end_tensor.data(); - float* cond_input = cond_input_tensor.data(); + void* cond_input; + bool cond_half = false; + if (cond_input_tensor.dtype() == at::kHalf) { + cond_input = (void*)cond_input_tensor.data_ptr(); + cond_half = true; + } else { + cond_input = (void*)cond_input_tensor.data(); + } float** in_layer_weights_prev = (float**)malloc(num_layers*sizeof(float*)); float** in_layer_weights_curr = (float**)malloc(num_layers*sizeof(float*)); @@ -84,7 +91,7 @@ int infer(at::Tensor samples_tensor, conv_out, conv_end, use_embed_tanh, - cond_input, + cond_input, cond_half, implementation, samples);