Resolve PR comments

This commit is contained in:
Swetha B S 2025-06-06 03:46:05 -07:00 committed by mcw
parent d6cc466bd4
commit 994e02a5eb
2 changed files with 22 additions and 22 deletions

View File

@ -6074,9 +6074,9 @@ template <typename BLOC_TYPE, int64_t INTER_SIZE, int64_t NB_COLS, ggml_type PAR
return false;
}
void forward_get_rows(const ggml_compute_params *params,
ggml_tensor *dst) {
const ggml_tensor *src0 = dst->src[0];
void forward_get_rows(const ggml_compute_params * params,
ggml_tensor * dst) {
const ggml_tensor * src0 = dst->src[0];
switch (src0->type) {
case GGML_TYPE_Q4_0: {
@ -6089,10 +6089,10 @@ template <typename BLOC_TYPE, int64_t INTER_SIZE, int64_t NB_COLS, ggml_type PAR
}
static void ggml_compute_forward_get_rows_q4_0x8(
const ggml_compute_params *params,
ggml_tensor *dst) {
const ggml_tensor *src0 = dst->src[0];
const ggml_tensor *src1 = dst->src[1];
const ggml_compute_params * params,
ggml_tensor * dst) {
const ggml_tensor * src0 = dst->src[0];
const ggml_tensor * src1 = dst->src[1];
GGML_TENSOR_BINARY_OP_LOCALS
@ -6132,10 +6132,10 @@ template <typename BLOC_TYPE, int64_t INTER_SIZE, int64_t NB_COLS, ggml_type PAR
int row_group_idx = i01 / nrows_interleaved;
const int row_idx_in_group = i01 % nrows_interleaved;
const char *base_ptr_for_higher_dims_in_src0 = (const char *)src0->data + i11 * nb02 + i12 * nb03;
const char * base_ptr_for_higher_dims_in_src0 = (const char *)src0->data + i11 * nb02 + i12 * nb03;
// Pointer to the first block_q4_0x8 of the identified row_group_idx
const block_q4_0x8 *p_first_repacked_block_of_group_x8 = (const block_q4_0x8 *)(base_ptr_for_higher_dims_in_src0 + row_group_idx * stride_between_actual_row_groups);
const block_q4_0x8 * p_first_repacked_block_of_group_x8 = (const block_q4_0x8 *)(base_ptr_for_higher_dims_in_src0 + row_group_idx * stride_between_actual_row_groups);
dequantize_row_q4_0x8(
p_first_repacked_block_of_group_x8,
@ -6152,8 +6152,8 @@ template <typename BLOC_TYPE, int64_t INTER_SIZE, int64_t NB_COLS, ggml_type PAR
* @param row_idx_in_group Index (0-7) of the logical row to dequantize.
*/
static void dequantize_row_q4_0x8(
const block_q4_0x8 *GGML_RESTRICT p_repacked_group_column_blocks,
float *GGML_RESTRICT y,
const block_q4_0x8 * GGML_RESTRICT p_repacked_group_column_blocks,
float * GGML_RESTRICT y,
int64_t k,
int row_idx_in_group) {
const int GGML_Q4_0_X8_INTERLEAVE_SIZE = 8;
@ -6168,23 +6168,23 @@ template <typename BLOC_TYPE, int64_t INTER_SIZE, int64_t NB_COLS, ggml_type PAR
const int qk4_0_half_elements = QK4_0 / 2;
for (int i = 0; i < nb; ++i) {
const block_q4_0x8 *current_column_repacked_block = &p_repacked_group_column_blocks[i];
const block_q4_0x8 * current_column_repacked_block = &p_repacked_group_column_blocks[i];
const float d_val = GGML_FP16_TO_FP32(current_column_repacked_block->d[row_idx_in_group]);
float *y_curr = y + i * QK4_0;
float * y_curr = y + i * QK4_0;
const int8_t *qs_first_half_repacked_ptr = &(current_column_repacked_block->qs[row_idx_in_group * bytes_for_half_elements]);
const int8_t * qs_first_half_repacked_ptr = &(current_column_repacked_block->qs[row_idx_in_group * bytes_for_half_elements]);
uint64_t first_half_chunk_u64;
memcpy(&first_half_chunk_u64, qs_first_half_repacked_ptr, sizeof(uint64_t));
first_half_chunk_u64 ^= xor_mask; // Reverse the XOR
const uint8_t *original_qs_first_half_bytes = (const uint8_t *)&first_half_chunk_u64;
const uint8_t * original_qs_first_half_bytes = (const uint8_t *)&first_half_chunk_u64;
const int8_t *qs_second_half_repacked_ptr = &(current_column_repacked_block->qs[offset_to_second_half_data + (row_idx_in_group * bytes_for_half_elements)]);
const int8_t * qs_second_half_repacked_ptr = &(current_column_repacked_block->qs[offset_to_second_half_data + (row_idx_in_group * bytes_for_half_elements)]);
uint64_t second_half_chunk_u64;
memcpy(&second_half_chunk_u64, qs_second_half_repacked_ptr, sizeof(uint64_t));
second_half_chunk_u64 ^= xor_mask; // Reverse the XOR
const uint8_t *original_qs_second_half_bytes = (const uint8_t *)&second_half_chunk_u64;
const uint8_t * original_qs_second_half_bytes = (const uint8_t *)&second_half_chunk_u64;
// dequantizing all QK4_0's for this block.
for (int j = 0; j < bytes_for_half_elements; ++j) {
@ -6530,10 +6530,10 @@ class extra_buffer_type : ggml::cpu::extra_buffer_type {
//if (op->src[1]->type == GGML_TYPE_Q8_0) {
// return true;
//}
} else if (op->op == GGML_OP_GET_ROWS
&& op->src[0]->buffer
&& (ggml_n_dims(op->src[0]) == 2)
&& op->src[0]->buffer->buft == ggml_backend_cpu_aarch64_buffer_type()
} else if (op->op == GGML_OP_GET_ROWS
&& op->src[0]->buffer
&& (ggml_n_dims(op->src[0]) == 2)
&& op->src[0]->buffer->buft == ggml_backend_cpu_aarch64_buffer_type()
&& ggml_aarch64_get_optimal_repack_type(op->src[0])) {
if (op->src[1]->buffer && !ggml_backend_buft_is_host(op->src[1]->buffer->buft)) {
return false;

View File

@ -1450,7 +1450,7 @@ static bool weight_buft_supported(const whisper_hparams & hparams, ggml_tensor *
ggml_context * ctx = ctx_ptr.get();
ggml_tensor * op_tensor = nullptr;
int64_t n_ctx = hparams.n_audio_ctx;
switch (op) {