Flatten batch to be checked
This commit is contained in:
@@ -47,11 +47,15 @@ public:
|
|||||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||||
dnnType *dstData = reinterpret_cast<dnnType*>(outputs[0]);
|
dnnType *dstData = reinterpret_cast<dnnType*>(outputs[0]);
|
||||||
checkCuda( cudaMemcpy(dstData, srcData, rows*cols*sizeof(dnnType), cudaMemcpyDeviceToDevice));
|
checkCuda( cudaMemcpyAsync(dstData, srcData, batchSize*rows*cols*sizeof(dnnType), cudaMemcpyDeviceToDevice, stream));
|
||||||
|
|
||||||
float const alpha(1.0);
|
checkERROR( cublasSetStream(handle, stream) );
|
||||||
float const beta(0.0);
|
for(int i=0; i<batchSize; i++) {
|
||||||
checkERROR( cublasSgeam( handle, CUBLAS_OP_T, CUBLAS_OP_N, rows, cols, &alpha, srcData, cols, &beta, srcData, rows, dstData, rows ));
|
float const alpha(1.0);
|
||||||
|
float const beta(0.0);
|
||||||
|
int offset = i*rows*cols;
|
||||||
|
checkERROR( cublasSgeam( handle, CUBLAS_OP_T, CUBLAS_OP_N, rows, cols, &alpha, srcData + offset, cols, &beta, srcData + offset, rows, dstData + offset, rows ));
|
||||||
|
}
|
||||||
return 0;
|
return 0;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -46,7 +46,8 @@ int main(int argc, char *argv[]) {
|
|||||||
TIMER_STOP
|
TIMER_STOP
|
||||||
|
|
||||||
// control output
|
// control output
|
||||||
for(int o=1; o<netRT.getBuffersN(); o++) {
|
std::cout<<"Output Buffers: "<<netRT.getBuffersN()-1<<"\n";
|
||||||
|
for(int o=1; o<netRT.getBuffersN(); o++) {
|
||||||
for(int b=1; b<BATCH_SIZE; b++) {
|
for(int b=1; b<BATCH_SIZE; b++) {
|
||||||
dnnType *out_d = (dnnType*) netRT.buffersRT[o];
|
dnnType *out_d = (dnnType*) netRT.buffersRT[o];
|
||||||
dnnType *out0_d = out_d;
|
dnnType *out0_d = out_d;
|
||||||
|
|||||||
Reference in New Issue
Block a user