src0_d + i3*(src0->nb[3] / sizeof(T)),
src1_d + i3*(src1->nb[3] / sizeof(T)),
dst_d + i3*( dst->nb[3] / sizeof(T)),
- src0->ne[0], src0->ne[1], src0->ne[2],
- dst->ne[0], dst->ne[1], dst->ne[2], dim, stream);
+ ggml_row_size(src0->type, src0->ne[0])/sizeof(T), src0->ne[1], src0->ne[2],
+ ggml_row_size(dst->type, dst->ne[0])/sizeof(T), dst->ne[1], dst->ne[2], dim, stream);
}
} else {
const size_t size0 = ggml_nbytes(src0);
CUDA_CHECK(cudaMemcpyAsync((char *) dst->data + size0, src1->data, size1, cudaMemcpyDeviceToDevice, stream));
}
} else {
+ GGML_ASSERT(!ggml_is_quantized(src0->type));
+
dim3 grid_dim(dst->ne[1], dst->ne[2], dst->ne[3]);
auto launch_kernel = [&](auto dim) {
concat_non_cont<T, dim><<<grid_dim, CUDA_CONCAT_BLOCK_SIZE, 0, stream>>>(
GGML_ASSERT(src0->type == src1->type);
GGML_ASSERT(dst->type == src0->type);
- GGML_ASSERT(!ggml_is_quantized(src0->type));
- GGML_ASSERT(ggml_blck_size(src0->type) == 1);
-
- switch (ggml_type_size(src0->type)) {
- case 1:
- concat_cuda<uint8_t>(src0, src1, dst, dim, stream);
- break;
- case 2:
- concat_cuda<uint16_t>(src0, src1, dst, dim, stream);
- break;
- case 4:
- concat_cuda<uint32_t>(src0, src1, dst, dim, stream);
- break;
- case 8:
- concat_cuda<uint64_t>(src0, src1, dst, dim, stream);
- break;
- default:
- GGML_ABORT("Unsupported type size: %zu", ggml_type_size(src0->type));
- break;
+
+ if (ggml_is_quantized(src0->type)) {
+ GGML_ASSERT(ggml_is_contiguous(src0));
+ GGML_ASSERT(ggml_is_contiguous(src1));
+ GGML_ASSERT(src0->ne[0] % ggml_blck_size(src0->type) == 0);
+ GGML_ASSERT(src1->ne[0] % ggml_blck_size(src1->type) == 0);
+
+ // if tensors are contiguous and ne[0] is multiple of the block size we can concat both tensors as byte tensors
+ concat_cuda<uint8_t>(src0, src1, dst, dim, stream);
+ } else {
+ GGML_ASSERT(ggml_blck_size(src0->type) == 1);
+
+ switch (ggml_type_size(src0->type)) {
+ case 1:
+ concat_cuda<uint8_t>(src0, src1, dst, dim, stream);
+ break;
+ case 2:
+ concat_cuda<uint16_t>(src0, src1, dst, dim, stream);
+ break;
+ case 4:
+ concat_cuda<uint32_t>(src0, src1, dst, dim, stream);
+ break;
+ case 8:
+ concat_cuda<uint64_t>(src0, src1, dst, dim, stream);
+ break;
+ default:
+ GGML_ABORT("Unsupported type size: %zu", ggml_type_size(src0->type));
+ break;
+ }
}
}
ggml_type src1_type = op->src[1]->type;
return src0_type == src1_type &&
src0_type == op->type &&
- !ggml_is_quantized(src0_type) &&
- ggml_blck_size(src0_type) == 1 &&
- (ggml_type_size(src0_type) == 1 ||
- ggml_type_size(src0_type) == 2 ||
- ggml_type_size(src0_type) == 4 ||
- ggml_type_size(src0_type) == 8);
+ (
+ (
+ ggml_is_quantized(src0_type) &&
+ ggml_is_contiguous(op->src[0]) &&
+ ggml_is_contiguous(op->src[1]) &&
+ op->src[0]->ne[0] % ggml_blck_size(src0_type) == 0 &&
+ op->src[1]->ne[0] % ggml_blck_size(src0_type) == 0
+ ) || (
+ !ggml_is_quantized(src0_type) &&
+ ggml_blck_size(src0_type) == 1 &&
+ (
+ ggml_type_size(src0_type) == 1 ||
+ ggml_type_size(src0_type) == 2 ||
+ ggml_type_size(src0_type) == 4 ||
+ ggml_type_size(src0_type) == 8
+ )
+ )
+ );
} break;
case GGML_OP_CONV_TRANSPOSE_1D:
{