#pragma clang diagnostic ignored "-Wunused-variable"
#pragma clang diagnostic ignored "-Wunused-function"
#pragma clang diagnostic ignored "-Wunused-but-set-variable"

#include <HAP_farf.h>
#include <HAP_perf.h>

#include <math.h>
#include <string.h>

#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "hex-common.h"
#include "dma-queue.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-tensor.h"
#include "hvx-utils.h"
#include "hvx-quant.h"
#include "matmul-ops.h"
#include "get-rows-ops.h"
#include "work-queue.h"

struct get_rows_context {
    struct htp_ops_context * octx;
    const struct htp_get_rows_kernel_params * kparams;
    struct htp_get_rows_vtcm_layout vtcm_layout;
    uint8_t * vtcm_base;
    uint32_t task_start;
    uint32_t tasks;
    uint32_t tasks_per_thread;
    uint32_t tile_size;
    uint32_t tile_stride;
    bool index_i32;
};

#define get_rows_preamble                      \
    const uint32_t ne00 = octx->src[0]->ne[0]; \
    const uint32_t ne01 = octx->src[0]->ne[1]; \
    const uint32_t ne02 = octx->src[0]->ne[2]; \
    const uint32_t ne03 = octx->src[0]->ne[3]; \
                                               \
    const uint32_t ne10 = octx->src[1]->ne[0]; \
    const uint32_t ne11 = octx->src[1]->ne[1]; \
    const uint32_t ne12 = octx->src[1]->ne[2]; \
    const uint32_t ne13 = octx->src[1]->ne[3]; \
                                               \
    const uint32_t ne0 = octx->dst->ne[0];     \
    const uint32_t ne1 = octx->dst->ne[1];     \
    const uint32_t ne2 = octx->dst->ne[2];     \
    const uint32_t ne3 = octx->dst->ne[3];     \
                                               \
    const uint32_t nb01 = octx->src[0]->nb[1]; \
    const uint32_t nb02 = octx->src[0]->nb[2]; \
    const uint32_t nb03 = octx->src[0]->nb[3]; \
                                               \
    const uint32_t nb10 = octx->src[1]->nb[0]; \
    const uint32_t nb11 = octx->src[1]->nb[1]; \
    const uint32_t nb12 = octx->src[1]->nb[2]; \
                                               \
    const uint32_t nb1 = octx->dst->nb[1];     \
    const uint32_t nb2 = octx->dst->nb[2];     \
    const uint32_t nb3 = octx->dst->nb[3];     \
                                               \
    const uint32_t nr = ne10 * ne11 * ne12;

#define GET_ROWS_THREAD_ST_FN(IDX_TYPE)                                                                                 \
static void get_rows_thread_st_##IDX_TYPE(unsigned int nth, unsigned int ith, void *data) {                             \
    struct get_rows_context * grctx = (struct get_rows_context *)data;                                                  \
    struct htp_ops_context * octx = grctx->octx;                                                                        \
    const struct htp_get_rows_kernel_params * kparams = grctx->kparams;                                                 \
    get_rows_preamble;                                                                                                  \
    const uint32_t dr  = grctx->tasks_per_thread;                                                                       \
    const uint32_t ir0 = grctx->task_start + dr * ith;                                                                  \
    if (ir0 >= grctx->task_start + grctx->tasks) {                                                                      \
        return;                                                                                                         \
    }                                                                                                                   \
    const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks);                                               \
    const uint32_t row_size_bytes = htp_tensor_get_row_size(octx->src[0]->type, ne00);                                  \
    dma_queue * dma_q = octx->ctx->dma[ith];                                                                            \
    for (uint32_t i = ir0; i < ir1; ++i) {                                                                              \
        const uint32_t i12 = fastdiv(i, &kparams->div_ne10_ne11);                                                       \
        const uint32_t rem = i - i12 * ne11 * ne10;                                                                     \
        const uint32_t i11 = fastdiv(rem, &kparams->div_ne10);                                                          \
        const uint32_t i10 = rem - i11 * ne10;                                                                          \
        const IDX_TYPE * src1_ptr = (const IDX_TYPE *)(uintptr_t)(octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12); \
        const uint32_t i01 = (uint32_t)*src1_ptr;                                                                       \
        assert(i01 < ne01);                                                                                             \
        const uint32_t q02 = fastdiv(i11, &kparams->div_ne02);                                                          \
        const uint32_t i02 = i11 - q02 * ne02;                                                                          \
        const uint32_t q03 = fastdiv(i12, &kparams->div_ne03);                                                          \
        const uint32_t i03 = i12 - q03 * ne03;                                                                          \
        const dma_addr_t src0_data = octx->src[0]->data + i01*nb01 + i02*nb02 + i03*nb03;                               \
        const dma_addr_t dst_data  = octx->dst->data    + i10*nb1  + i11*nb2  + i12*nb3;                                \
        while (!dma_queue_push(dma_q, dma_make_data(dst_data, src0_data), nb1, nb01,                                    \
                               row_size_bytes, 1)) {                                                                    \
            dma_queue_pop(dma_q);                                                                                       \
        }                                                                                                               \
    }                                                                                                                   \
    dma_queue_flush(dma_q);                                                                                             \
}

GET_ROWS_THREAD_ST_FN(int32_t)
GET_ROWS_THREAD_ST_FN(int64_t)

#define GET_ROWS_THREAD_DT_FN(TYPE_NAME, SRC0_SIZE_EXPR, IDX_TYPE, COMPUTE_EXPR)                                                \
static void get_rows_thread_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsigned int ith, void *data) {                          \
    struct get_rows_context * grctx = (struct get_rows_context *)data;                                                          \
    struct htp_ops_context * octx = grctx->octx;                                                                                \
    const struct htp_get_rows_kernel_params * kparams = grctx->kparams;                                                         \
    get_rows_preamble;                                                                                                          \
    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                                      \
    const uint32_t dr  = grctx->tasks_per_thread;                                                                               \
    const uint32_t ir0 = grctx->task_start + dr * ith;                                                                          \
    if (ir0 >= grctx->task_start + grctx->tasks) {                                                                              \
        return;                                                                                                                 \
    }                                                                                                                           \
    const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks);                                                       \
    const uint32_t chunks_per_row = kparams->chunks_per_row;                                                                    \
    const uint32_t chunk_size     = kparams->chunk_size;                                                                        \
    dma_queue * dma_q = octx->ctx->dma[ith];                                                                                    \
    const struct htp_get_rows_vtcm_layout * vtcm_layout = &grctx->vtcm_layout;                                                  \
    uint8_t * vtcm_src0 = grctx->vtcm_base + vtcm_layout->off_src0 + ith * vtcm_layout->src0_bytes_per_thread;                  \
    uint8_t * vtcm_dst  = grctx->vtcm_base + vtcm_layout->off_dst  + ith * vtcm_layout->dst_bytes_per_thread;                   \
    for (uint32_t step = 0, spad_idx = 0; step < ir1 - ir0 && spad_idx < 2; ++step, spad_idx++) {                               \
        const uint32_t i = ir0 + step;                                                                                          \
        const uint32_t row_idx   = fastdiv(i, &kparams->div_chunks_per_row);                                                    \
        const uint32_t chunk_idx = i - row_idx * chunks_per_row;                                                                \
        const uint32_t i12 = fastdiv(row_idx, &kparams->div_ne10_ne11);                                                         \
        const uint32_t rem = row_idx - i12 * ne11 * ne10;                                                                       \
        const uint32_t i11 = fastdiv(rem, &kparams->div_ne10);                                                                  \
        const uint32_t i10 = rem - i11 * ne10;                                                                                  \
        const IDX_TYPE * src1_ptr = (const IDX_TYPE *)(uintptr_t)(octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12);         \
        const uint32_t i01 = (uint32_t)*src1_ptr;                                                                               \
        assert(i01 < ne01);                                                                                                     \
        const uint32_t q02 = fastdiv(i11, &kparams->div_ne02);                                                                  \
        const uint32_t i02 = i11 - q02 * ne02;                                                                                  \
        const uint32_t q03 = fastdiv(i12, &kparams->div_ne03);                                                                  \
        const uint32_t i03 = i12 - q03 * ne03;                                                                                  \
        const uint32_t offset = chunk_idx * chunk_size;                                                                         \
        const uint32_t cur_elems = (offset < ne00) ? MIN(chunk_size, ne00 - offset) : 0;                                        \
        const uint32_t cur_src0_bytes = SRC0_SIZE_EXPR(cur_elems);                                                              \
        const uint32_t cur_dst_bytes  = cur_elems * sizeof(float);                                                              \
        const dma_addr_t src0_data = octx->src[0]->data + i01*nb01 + i02*nb02 + i03*nb03 + SRC0_SIZE_EXPR(offset);              \
        dma_queue_push(dma_q,                                                                                                   \
                       dma_make_data(octx->dst->data,                                                                           \
                                     vtcm_dst + spad_idx * vtcm_layout->dst_spad_half_size),                                    \
                       cur_dst_bytes, vtcm_layout->dst_spad_half_size, cur_dst_bytes, 0);                                       \
        dma_queue_push(dma_q,                                                                                                   \
                       dma_make_data(vtcm_src0 + spad_idx * vtcm_layout->src0_spad_half_size, src0_data),                       \
                       vtcm_layout->src0_spad_half_size, cur_src0_bytes, cur_src0_bytes, 1);                                    \
    }                                                                                                                           \
    for (uint32_t step = 0; step < ir1 - ir0; ++step) {                                                                         \
        const uint32_t i = ir0 + step;                                                                                          \
        void * dst_spad = (void *) dma_queue_pop(dma_q).src;                                                                    \
        void * src_spad = (void *) dma_queue_pop(dma_q).dst;                                                                    \
        const uint32_t row_idx   = fastdiv(i, &kparams->div_chunks_per_row);                                                    \
        const uint32_t chunk_idx = i - row_idx * chunks_per_row;                                                                \
        const uint32_t i12 = fastdiv(row_idx, &kparams->div_ne10_ne11);                                                         \
        const uint32_t rem = row_idx - i12 * ne11 * ne10;                                                                       \
        const uint32_t i11 = fastdiv(rem, &kparams->div_ne10);                                                                  \
        const uint32_t i10 = rem - i11 * ne10;                                                                                  \
        const uint32_t offset = chunk_idx * chunk_size;                                                                         \
        const uint32_t cur_elems = (offset < ne00) ? MIN(chunk_size, ne00 - offset) : 0;                                        \
        const uint32_t cur_dst_bytes  = cur_elems * sizeof(float);                                                              \
        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, i);                                                                   \
        COMPUTE_EXPR;                                                                                                           \
        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, i);                                                                    \
        const dma_addr_t dst_data  = octx->dst->data + i10*nb1 + i11*nb2 + i12*nb3 + offset * sizeof(float);                    \
        dma_queue_push(dma_q,                                                                                                   \
                       dma_make_data(dst_data, dst_spad),                                                                       \
                       cur_dst_bytes, vtcm_layout->dst_spad_half_size, cur_dst_bytes, 1);                                       \
        const uint32_t next_step = step + 2;                                                                                    \
        if (next_step < ir1 - ir0) {                                                                                            \
            const uint32_t pi = ir0 + next_step;                                                                                \
            const uint32_t prow_idx   = fastdiv(pi, &kparams->div_chunks_per_row);                                              \
            const uint32_t pchunk_idx = pi - prow_idx * chunks_per_row;                                                         \
            const uint32_t pi12 = fastdiv(prow_idx, &kparams->div_ne10_ne11);                                                   \
            const uint32_t prem = prow_idx - pi12 * ne11 * ne10;                                                                \
            const uint32_t pi11 = fastdiv(prem, &kparams->div_ne10);                                                            \
            const uint32_t pi10 = prem - pi11 * ne10;                                                                           \
            const IDX_TYPE * psrc1_ptr = (const IDX_TYPE *)(uintptr_t)(octx->src[1]->data + pi10*nb10 + pi11*nb11 + pi12*nb12); \
            const uint32_t pi01 = (uint32_t)*psrc1_ptr;                                                                         \
            assert(pi01 < ne01);                                                                                                \
            const uint32_t pq02 = fastdiv(pi11, &kparams->div_ne02);                                                            \
            const uint32_t pi02 = pi11 - pq02 * ne02;                                                                           \
            const uint32_t pq03 = fastdiv(pi12, &kparams->div_ne03);                                                            \
            const uint32_t pi03 = pi12 - pq03 * ne03;                                                                           \
            const uint32_t poffset = pchunk_idx * chunk_size;                                                                   \
            const uint32_t pcur_elems = (poffset < ne00) ? MIN(chunk_size, ne00 - poffset) : 0;                                 \
            const uint32_t pcur_src0_bytes = SRC0_SIZE_EXPR(pcur_elems);                                                        \
            const dma_addr_t psrc0_data =                                                                                       \
                octx->src[0]->data + pi01*nb01 + pi02*nb02 + pi03*nb03 + SRC0_SIZE_EXPR(poffset);                               \
            dma_queue_push(dma_q,                                                                                               \
                           dma_make_data(src_spad, psrc0_data),                                                                 \
                           vtcm_layout->src0_spad_half_size, pcur_src0_bytes, pcur_src0_bytes, 1);                              \
        }                                                                                                                       \
    }                                                                                                                           \
    dma_queue_flush(dma_q);                                                                                                     \
}

#define F16_BYTES(n)  ((n) * sizeof(__fp16))
#define Q8_0_BYTES(n) (((n) / 32) * sizeof(block_q8_0))

static __attribute__((noinline)) void compute_get_rows_f16(float * dst_spad, const void * src_spad, uint32_t cur_elems) {
    hvx_dequantize_row_f16_f32(dst_spad, src_spad, cur_elems);
}

static __attribute__((noinline)) void compute_get_rows_q8_0(float * dst_spad, const void * src_spad, uint32_t cur_elems) {
    hvx_dequantize_row_q8_0_f32(dst_spad, src_spad, cur_elems);
}

GET_ROWS_THREAD_DT_FN(f16,  F16_BYTES,  int32_t, { compute_get_rows_f16((float *)dst_spad, src_spad, cur_elems); })
GET_ROWS_THREAD_DT_FN(f16,  F16_BYTES,  int64_t, { compute_get_rows_f16((float *)dst_spad, src_spad, cur_elems); })

GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int32_t, { compute_get_rows_q8_0((float *)dst_spad, src_spad, cur_elems); })
GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int64_t, { compute_get_rows_q8_0((float *)dst_spad, src_spad, cur_elems); })


static __attribute__((noinline)) void compute_get_rows_tiled(float * dst, const uint8_t * tile, uint32_t row, bool q4) {
    const HVX_VectorPred first2 = Q6_Q_vsetq_R(2);
    const HVX_VectorPred first4 = Q6_Q_vsetq_R(4);
    HVX_Vector vq = Q6_V_vzero();
    if (q4) {
        const HVX_VectorPred first1 = Q6_Q_vsetq_R(1);
        const HVX_VectorPred first3 = Q6_Q_vsetq_R(3);
        for (int group = 3; group >= 0; --group) {
            const HVX_Vector v = Q6_V_vror_VR(hvx_vmem(tile + group * VLEN), row);
            // Four planes contribute bytes at 0, 32, 64 and 96 after rotation.
            HVX_Vector packed = Q6_V_vmux_QVV(first1, v, Q6_V_vror_VR(v, 31));
            packed = Q6_V_vmux_QVV(first2, packed, Q6_V_vror_VR(v, 62));
            packed = Q6_V_vmux_QVV(first3, packed, Q6_V_vror_VR(v, 93));
            vq = Q6_V_vmux_QVV(first4, packed, Q6_V_vror_VR(vq, VLEN - 4));
        }
        const HVX_Vector lo = Q6_V_vand_VV(vq, Q6_Vb_vsplat_R(0x0F));
        const HVX_Vector hi = Q6_Vub_vlsr_VubR(vq, 4);
        vq = Q6_V_lo_W(Q6_W_vshuff_VVR(hi, lo, -1));
        vq = Q6_Vb_vsub_VbVb(vq, Q6_Vb_vsplat_R(8));
    } else {
        for (int group = 7; group >= 0; --group) {
            const HVX_Vector v = Q6_V_vror_VR(hvx_vmem(tile + group * VLEN), 2 * row);
            // Two planes contribute halfwords at 0 and 64 after rotation.
            const HVX_Vector packed = Q6_V_vmux_QVV(first2, v, Q6_V_vror_VR(v, 62));
            vq = Q6_V_vmux_QVV(first4, packed, Q6_V_vror_VR(vq, VLEN - 4));
        }
    }
    const HVX_Vector scales = hvx_vmem(tile + (q4 ? 512 : 1024));
    const HVX_Vector scale_hf = hvx_vec_repl_f16(Q6_V_vror_VR(scales, 2 * row));
    const HVX_Vector scale = Q6_V_lo_W(hvx_vec_f16_to_f32(scale_hf));
    const HVX_VectorPair p16 = Q6_Wh_vunpack_Vb(vq);
    const HVX_VectorPair p32 = Q6_Ww_vunpack_Vh(Q6_V_lo_W(p16));
    const HVX_Vector values = hvx_vec_mul_f32_f32(Q6_Vsf_equals_Vw(Q6_V_lo_W(p32)), scale);
    *(HVX_Vector *) dst = values;
}

struct get_rows_tiled_task {
    dma_addr_t tile_src_base;
    dma_addr_t dst_data;
    uint32_t   row;
};

static inline struct get_rows_tiled_task get_rows_tiled_calc_task(
    const struct htp_ops_context * octx,
    const struct get_rows_context * grctx,
    uint32_t i,
    uint32_t n_k_tiles,
    uint32_t tile_size
) {
    const struct htp_get_rows_kernel_params * kparams = grctx->kparams;
    get_rows_preamble;

    const uint32_t i12 = fastdiv(i, &kparams->div_ne10_ne11);
    const uint32_t rem = i - i12 * ne11 * ne10;
    const uint32_t i11 = fastdiv(rem, &kparams->div_ne10);
    const uint32_t i10 = rem - i11 * ne10;
    const dma_addr_t src1_data = octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12;
    const uint32_t i01 = grctx->index_i32 ? *(const int32_t *)(uintptr_t) src1_data : (uint32_t) *(const int64_t *)(uintptr_t) src1_data;
    assert(i01 < ne01);

    const uint32_t q02 = fastdiv(i11, &kparams->div_ne02);
    const uint32_t i02 = i11 - q02 * ne02;
    const uint32_t q03 = fastdiv(i12, &kparams->div_ne03);
    const uint32_t i03 = i12 - q03 * ne03;
    const uint32_t column_tile = i01 / HTP_MM_HMX_TILE_N_ROWS;
    const uint32_t row = i01 % HTP_MM_HMX_TILE_N_ROWS;
    const dma_addr_t matrix = octx->src[0]->data + i02*nb02 + i03*nb03;

    struct get_rows_tiled_task task;
    task.tile_src_base = matrix + (column_tile * n_k_tiles) * tile_size;
    task.dst_data      = octx->dst->data + i10*nb1 + i11*nb2 + i12*nb3;
    task.row           = row;
    return task;
}

static void get_rows_thread_tiled(unsigned int nth, unsigned int ith, void * data) {
    struct get_rows_context * grctx = (struct get_rows_context *) data;
    struct htp_ops_context * octx = grctx->octx;
    const struct htp_get_rows_kernel_params * kparams = grctx->kparams;
    get_rows_preamble;

    const uint32_t dr  = grctx->tasks_per_thread;
    const uint32_t ir0 = grctx->task_start + dr * ith;
    if (ir0 >= grctx->task_start + grctx->tasks) {
        return;
    }

    const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks);
    const uint32_t n_k_tiles = ne00 / HTP_MM_HMX_TILE_N_COLS;
    const struct htp_get_rows_vtcm_layout * vtcm_layout = &grctx->vtcm_layout;
    uint8_t * src_spad_base = grctx->vtcm_base + vtcm_layout->off_src0 + ith * vtcm_layout->src0_bytes_per_thread;
    uint8_t * dst_spad_base = grctx->vtcm_base + vtcm_layout->off_dst + ith * vtcm_layout->dst_bytes_per_thread;
    dma_queue * dma_q = octx->ctx->dma[ith];
    struct htp_thread_trace * tr = &octx->ctx->trace[ith];

    const uint32_t tile_size   = grctx->tile_size;
    const uint32_t tile_stride = grctx->tile_stride;
    const uint32_t dst_bytes   = ne00 * sizeof(float);
    const bool is_q4 = (octx->src[0]->type == HTP_TYPE_Q4_0);

    for (uint32_t step = 0, spad_idx = 0; step < ir1 - ir0 && spad_idx < 2; ++step, ++spad_idx) {
        const uint32_t i = ir0 + step;
        struct get_rows_tiled_task task = get_rows_tiled_calc_task(octx, grctx, i, n_k_tiles, tile_size);

        // Dummy writeback to prime the queue with dst descriptor
        dma_queue_push(dma_q,
                       dma_make_data(task.dst_data, dst_spad_base + spad_idx * vtcm_layout->dst_spad_half_size),
                       dst_bytes, vtcm_layout->dst_spad_half_size, dst_bytes, 0);

        // Prefetch row tiles
        dma_queue_push(dma_q,
                       dma_make_data(src_spad_base + spad_idx * vtcm_layout->src0_spad_half_size, task.tile_src_base),
                       tile_stride, tile_size, tile_size, n_k_tiles);
    }

    for (uint32_t step = 0; step < ir1 - ir0; ++step) {
        const uint32_t i = ir0 + step;
        float * dst_spad   = (float *) dma_queue_pop(dma_q).src;
        uint8_t * src_spad = (uint8_t *) dma_queue_pop(dma_q).dst;

        struct get_rows_tiled_task task = get_rows_tiled_calc_task(octx, grctx, i, n_k_tiles, tile_size);

        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
        for (uint32_t k_tile = 0; k_tile < n_k_tiles; ++k_tile) {
            const uint8_t * tile = src_spad + k_tile * tile_stride;
            float * dst_block = dst_spad + k_tile * HTP_MM_HMX_TILE_N_COLS;
            compute_get_rows_tiled(dst_block, tile, task.row, is_q4);
        }
        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);

        // Real writeback of dst_spad
        dma_queue_push(dma_q,
                       dma_make_data(task.dst_data, dst_spad),
                       dst_bytes, vtcm_layout->dst_spad_half_size, dst_bytes, 1);

        const uint32_t next_step = step + 2;
        if (next_step < ir1 - ir0) {
            const uint32_t ni = ir0 + next_step;
            struct get_rows_tiled_task next_task = get_rows_tiled_calc_task(octx, grctx, ni, n_k_tiles, tile_size);
            dma_queue_push(dma_q,
                           dma_make_data(src_spad, next_task.tile_src_base),
                           tile_stride, tile_size, tile_size, n_k_tiles);
        }
    }

    dma_queue_flush(dma_q);
}

int op_get_rows(struct htp_ops_context * octx) {
    const struct htp_get_rows_kernel_params * kparams = (const struct htp_get_rows_kernel_params *) octx->kernel_params;

    if (octx->src[0]->type != HTP_TYPE_F32 &&
         octx->src[0]->type != HTP_TYPE_F16 &&
         octx->src[0]->type != HTP_TYPE_Q4_0 &&
         octx->src[0]->type != HTP_TYPE_Q8_0 &&
         octx->src[0]->type != HTP_TYPE_I32) {
        return HTP_STATUS_NO_SUPPORT;
    }

    if (kparams->kernel_type == HTP_GET_ROWS_KERNEL_SAMETYPE) {
        if (octx->src[0]->type != octx->dst->type) {
            return HTP_STATUS_NO_SUPPORT;
        }
    } else {
        if (octx->dst->type != HTP_TYPE_F32) {
            return HTP_STATUS_NO_SUPPORT;
        }
    }

    if (octx->src[1]->type != HTP_TYPE_I32 && octx->src[1]->type != HTP_TYPE_I64) {
        return HTP_STATUS_NO_SUPPORT;
    }

    if (htp_tensor_is_extended(octx->src[1])) {
        return HTP_STATUS_NO_SUPPORT;
    }

    const struct htp_tensor * dst = octx->dst;
    const uint32_t total_tasks    = kparams->total_tasks;
    const size_t dst_row_size     = htp_tensor_get_row_size(dst->type, dst->ne[0]);

    uint32_t task_start = 0;
    uint32_t tasks      = total_tasks;

    if (octx->ctx->mdev.count > 1) {
        uint32_t tasks_per_chunk = 1;
        htp_tensor_mdev_rows_per_chunk(dst, dst_row_size / dst->ne[0], (uint32_t) dst_row_size, &tasks_per_chunk);
        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_tasks, tasks_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
        task_start = range.start;
        tasks      = range.count;
    }

    if (tasks == 0) {
        return HTP_STATUS_OK;
    }

    if (!htp_ops_context_set_n_threads(octx, (uint32_t) kparams->n_threads)) {
        return HTP_STATUS_INVAL_PARAMS;
    }

    const uint32_t n_threads = octx->n_threads;

    struct get_rows_context grctx;
    grctx.octx = octx;
    grctx.kparams = kparams;
    grctx.vtcm_base = (uint8_t *)octx->ctx->vtcm_base;
    grctx.task_start = task_start;
    grctx.tasks = tasks;
    grctx.tasks_per_thread = octx->ctx->mdev.count == 1 ? kparams->tasks_per_thread : fastdiv(tasks + n_threads - 1, &octx->n_threads_div);
    grctx.tile_size = octx->src[0]->type == HTP_TYPE_Q4_0 ? HTP_MM_WEIGHT_TILE_SIZE_Q4_0 : HTP_MM_WEIGHT_TILE_SIZE_Q8_0;
    grctx.tile_stride = (grctx.tile_size + 127) & ~127;
    grctx.index_i32 = octx->src[1]->type == HTP_TYPE_I32;

    const uint32_t ne00 = octx->src[0]->ne[0];
    htp_get_rows_vtcm_layout_build(&grctx.vtcm_layout, kparams->kernel_type, octx->src[0]->type, ne00, n_threads);

    if (grctx.vtcm_layout.total_bytes > octx->ctx->vtcm_size) {
        FARF(ERROR, "get-rows: VTCM reservation %zu is too small, needed %zu\n",
             octx->ctx->vtcm_size, grctx.vtcm_layout.total_bytes);
        return HTP_STATUS_INVAL_PARAMS;
    }

    const bool is_i32 = (octx->src[1]->type == HTP_TYPE_I32);

    work_queue_func_t q_func = NULL;
    switch (kparams->kernel_type) {
        case HTP_GET_ROWS_KERNEL_SAMETYPE:
            q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_st_int32_t : get_rows_thread_st_int64_t);
            break;
        case HTP_GET_ROWS_KERNEL_TILED:
            q_func = get_rows_thread_tiled;
            break;
        case HTP_GET_ROWS_KERNEL_FLAT:
            switch (octx->src[0]->type) {
                case HTP_TYPE_F16:  q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_f16_int32_t  : get_rows_thread_f16_int64_t);  break;
                case HTP_TYPE_Q8_0: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_q8_0_int32_t : get_rows_thread_q8_0_int64_t); break;
                default:            return HTP_STATUS_NO_SUPPORT;
            }
            break;
        default:
            return HTP_STATUS_NO_SUPPORT;
    }

    FARF(HIGH, "get-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu kernel-type %d n-threads %d\n",
         octx->src[0]->ne[0], octx->src[0]->ne[1], octx->src[0]->ne[2], octx->src[0]->ne[3],
         octx->src[1]->ne[0], octx->src[1]->ne[1], octx->src[1]->ne[2], octx->src[1]->ne[3],
         octx->dst->ne[0], octx->dst->ne[1], octx->dst->ne[2], octx->dst->ne[3],
         grctx.vtcm_layout.src0_bytes_per_thread * n_threads,
         grctx.vtcm_layout.dst_bytes_per_thread  * n_threads,
         kparams->kernel_type, n_threads);

    work_queue_run(octx->ctx->work_queue, q_func, &grctx, n_threads);
    return HTP_STATUS_OK;
}
