@@ -71,6 +71,11 @@ namespace Inference {
7171 }
7272 }
7373
74+ void FloatTransformer::clear_kv_cache () {
75+ std::memset (state.key_cache .data (), 0 , state.key_cache .size () * sizeof (float ));
76+ std::memset (state.value_cache .data (), 0 , state.value_cache .size () * sizeof (float ));
77+ }
78+
7479 void FloatTransformer::rmsnorm (float *__restrict__ o, const float *__restrict__ x, const float *__restrict__ weight,
7580 const int size) {
7681 float ss = 0 .0f ;
@@ -209,13 +214,24 @@ namespace Inference {
209214 }
210215
211216 void FloatTransformer::softmax (float *x, const int size) {
212- const float max_val = *std::max_element (x, x + size);
217+ // Find max for numerical stability
218+ float max_val = x[0 ];
219+ #pragma omp parallel for reduction(max:max_val)
220+ for (int i = 1 ; i < size; i++) {
221+ if (x[i] > max_val) max_val = x[i];
222+ }
223+
224+ // Exp and sum
213225 float sum = 0 .0f ;
226+ #pragma omp parallel for reduction(+:sum)
214227 for (int i = 0 ; i < size; i++) {
215228 x[i] = std::exp (x[i] - max_val);
216229 sum += x[i];
217230 }
231+
232+ // Normalize
218233 const float inv_sum = 1 .0f / sum;
234+ #pragma omp parallel for
219235 for (int i = 0 ; i < size; i++) {
220236 x[i] *= inv_sum;
221237 }
@@ -477,19 +493,30 @@ namespace Inference {
477493 }
478494 }
479495
496+ // Precompute attention scale factor (moved out of inner loops for efficiency)
497+ const float att_scale = 1 .0f / std::sqrt (static_cast <float >(head_size));
498+
480499 int h;
481500#pragma omp parallel for private(h)
482501 for (h = 0 ; h < p->n_heads ; h++) {
483502 const float *q = s->q .data () + h * head_size;
484503 float *att = s->att .data () + h * p->seq_len ;
485504 int t = 0 ;
486505
487- // Unrolled loop for t
506+ // Unrolled loop for t with prefetching
488507 for (; t <= pos - 4 ; t += 4 ) {
489508 const float *k0 = s->key_cache .data () + loff + t * kv_dim + (h / kv_mul) * head_size;
490509 const float *k1 = s->key_cache .data () + loff + (t + 1 ) * kv_dim + (h / kv_mul) * head_size;
491510 const float *k2 = s->key_cache .data () + loff + (t + 2 ) * kv_dim + (h / kv_mul) * head_size;
492511 const float *k3 = s->key_cache .data () + loff + (t + 3 ) * kv_dim + (h / kv_mul) * head_size;
512+
513+ // Prefetch next batch of key vectors
514+ if (t + 4 <= pos - 4 ) {
515+ __builtin_prefetch (s->key_cache .data () + loff + (t + 4 ) * kv_dim + (h / kv_mul) * head_size, 0 , 0 );
516+ __builtin_prefetch (s->key_cache .data () + loff + (t + 5 ) * kv_dim + (h / kv_mul) * head_size, 0 , 0 );
517+ __builtin_prefetch (s->key_cache .data () + loff + (t + 6 ) * kv_dim + (h / kv_mul) * head_size, 0 , 0 );
518+ __builtin_prefetch (s->key_cache .data () + loff + (t + 7 ) * kv_dim + (h / kv_mul) * head_size, 0 , 0 );
519+ }
493520
494521 float s0 = 0 .0f ;
495522 float s1 = 0 .0f ;
@@ -558,11 +585,10 @@ namespace Inference {
558585 s3 += qv * k3[i];
559586 }
560587
561- float scale = 1 .0f / std::sqrt (static_cast <float >(head_size));
562- att[t] = s0 * scale;
563- att[t + 1 ] = s1 * scale;
564- att[t + 2 ] = s2 * scale;
565- att[t + 3 ] = s3 * scale;
588+ att[t] = s0 * att_scale;
589+ att[t + 1 ] = s1 * att_scale;
590+ att[t + 2 ] = s2 * att_scale;
591+ att[t + 3 ] = s3 * att_scale;
566592 }
567593
568594 // Tail loop
@@ -600,14 +626,13 @@ namespace Inference {
600626 for (; i < head_size; i++) {
601627 score += q[i] * k_ptr[i];
602628 }
603- score /= std::sqrt (static_cast <float >(head_size));
604- att[t] = score;
629+ att[t] = score * att_scale;
605630 }
606631
607632 softmax (att, pos + 1 );
608633
609634 float *xb = s->xb .data () + h * head_size;
610- std::fill_n (xb, head_size, 0 . 0f );
635+ std::memset (xb, 0 , head_size * sizeof ( float ) );
611636
612637 for (int t = 0 ; t <= pos; t++) {
613638 const float *v_ptr = s->value_cache .data () + loff + t * kv_dim + (h / kv_mul) * head_size;
@@ -638,7 +663,22 @@ namespace Inference {
638663
639664 matmul (s->xb2 .data (), s->xb .data (), w->wo + l * dim * dim, dim, dim);
640665
641- for (int i = 0 ; i < dim; i++) {
666+ // Vectorized residual connection
667+ int i = 0 ;
668+ #if defined(__ARM_NEON)
669+ for (; i <= dim - 4 ; i += 4 ) {
670+ float32x4_t x_vec = vld1q_f32 (x + i);
671+ float32x4_t xb2_vec = vld1q_f32 (s->xb2 .data () + i);
672+ vst1q_f32 (x + i, vaddq_f32 (x_vec, xb2_vec));
673+ }
674+ #elif defined(__AVX2__)
675+ for (; i <= dim - 8 ; i += 8 ) {
676+ __m256 x_vec = _mm256_loadu_ps (x + i);
677+ __m256 xb2_vec = _mm256_loadu_ps (s->xb2 .data () + i);
678+ _mm256_storeu_ps (x + i, _mm256_add_ps (x_vec, xb2_vec));
679+ }
680+ #endif
681+ for (; i < dim; i++) {
642682 x[i] += s->xb2 [i];
643683 }
644684
@@ -647,17 +687,63 @@ namespace Inference {
647687 matmul (s->hb .data (), s->xb .data (), w->w1 + l * dim * hidden_dim, dim, hidden_dim);
648688 matmul (s->hb2 .data (), s->xb .data (), w->w3 + l * dim * hidden_dim, dim, hidden_dim);
649689
650- #pragma omp parallel for simd
651- for (int i = 0 ; i < hidden_dim; i++) {
652- float val = s->hb [i];
653- val *= (1 .0f / (1 .0f + std::exp (-val))); // Silu
654- val *= s->hb2 [i];
655- s->hb [i] = val;
690+ // SiLU activation with explicit SIMD
691+ int ii = 0 ;
692+ #if defined(__ARM_NEON)
693+ for (; ii <= hidden_dim - 4 ; ii += 4 ) {
694+ float32x4_t val = vld1q_f32 (s->hb .data () + ii);
695+ float32x4_t hb2_vec = vld1q_f32 (s->hb2 .data () + ii);
696+ // Compute sigmoid: 1 / (1 + exp(-val))
697+ // Using approximation for NEON
698+ float tmp[4 ];
699+ vst1q_f32 (tmp, val);
700+ tmp[0 ] = 1 .0f / (1 .0f + std::exp (-tmp[0 ]));
701+ tmp[1 ] = 1 .0f / (1 .0f + std::exp (-tmp[1 ]));
702+ tmp[2 ] = 1 .0f / (1 .0f + std::exp (-tmp[2 ]));
703+ tmp[3 ] = 1 .0f / (1 .0f + std::exp (-tmp[3 ]));
704+ float32x4_t sigmoid = vld1q_f32 (tmp);
705+ float32x4_t result = vmulq_f32 (vmulq_f32 (val, sigmoid), hb2_vec);
706+ vst1q_f32 (s->hb .data () + ii, result);
707+ }
708+ #elif defined(__AVX2__)
709+ for (; ii <= hidden_dim - 8 ; ii += 8 ) {
710+ __m256 val = _mm256_loadu_ps (s->hb .data () + ii);
711+ __m256 hb2_vec = _mm256_loadu_ps (s->hb2 .data () + ii);
712+ // Scalar sigmoid computation (exp is not vectorized in AVX2 without libmvec)
713+ float tmp[8 ], sig[8 ];
714+ _mm256_storeu_ps (tmp, val);
715+ for (int k = 0 ; k < 8 ; k++) {
716+ sig[k] = 1 .0f / (1 .0f + std::exp (-tmp[k]));
717+ }
718+ __m256 sigmoid = _mm256_loadu_ps (sig);
719+ __m256 result = _mm256_mul_ps (_mm256_mul_ps (val, sigmoid), hb2_vec);
720+ _mm256_storeu_ps (s->hb .data () + ii, result);
721+ }
722+ #endif
723+ for (; ii < hidden_dim; ii++) {
724+ float val = s->hb [ii];
725+ val *= (1 .0f / (1 .0f + std::exp (-val))); // SiLU
726+ val *= s->hb2 [ii];
727+ s->hb [ii] = val;
656728 }
657-
658729 matmul (s->xb .data (), s->hb .data (), w->w2 + l * dim * hidden_dim, hidden_dim, dim);
659730
660- for (int i = 0 ; i < dim; i++) {
731+ // Vectorized residual connection
732+ i = 0 ;
733+ #if defined(__ARM_NEON)
734+ for (; i <= dim - 4 ; i += 4 ) {
735+ float32x4_t x_vec = vld1q_f32 (x + i);
736+ float32x4_t xb_vec = vld1q_f32 (s->xb .data () + i);
737+ vst1q_f32 (x + i, vaddq_f32 (x_vec, xb_vec));
738+ }
739+ #elif defined(__AVX2__)
740+ for (; i <= dim - 8 ; i += 8 ) {
741+ __m256 x_vec = _mm256_loadu_ps (x + i);
742+ __m256 xb_vec = _mm256_loadu_ps (s->xb .data () + i);
743+ _mm256_storeu_ps (x + i, _mm256_add_ps (x_vec, xb_vec));
744+ }
745+ #endif
746+ for (; i < dim; i++) {
661747 x[i] += s->xb [i];
662748 }
663749 }
@@ -717,19 +803,125 @@ namespace Inference {
717803 const int dim = config.dim ;
718804 PacketHeader header{};
719805
720- const int start_layer = dist_config.split_layer ;
806+ int start_layer = dist_config.split_layer ;
721807
722808 const size_t packet_size = sizeof (PacketHeader) + dim * sizeof (float );
723809 if (transfer_buffer.size () < packet_size) transfer_buffer.resize (packet_size);
810+
811+ // Set packet size on transport for properly padded control messages
812+ dist_config.transport ->set_packet_size (packet_size);
724813
725814 std::cout << " Worker started. Processing layers " << start_layer << " to " << dist_config.end_layer - 1 <<
726815 std::endl;
727816
728817 while (true ) {
729818 try {
819+ // Check for control messages (non-blocking)
820+ ControlMessage ctrl_msg{};
821+ if (dist_config.transport ->recv_control_nonblocking (ctrl_msg)) {
822+ if (ctrl_msg.type == ControlMessageType::RESIZE_LAYERS ) {
823+ update_layer_config (ctrl_msg);
824+ start_layer = dist_config.split_layer ;
825+ std::cout << " Worker: Layer config updated. Now processing layers "
826+ << start_layer << " to " << dist_config.end_layer - 1 << std::endl;
827+
828+ // Send ACK back (optional, for confirmation)
829+ ControlMessage ack{};
830+ ack.type = ControlMessageType::ACK ;
831+ ack.split_layer = start_layer;
832+ ack.end_layer = dist_config.end_layer ;
833+ ack.is_tail = dist_config.is_tail ;
834+ dist_config.transport ->send_control (ack);
835+ } else if (ctrl_msg.type == ControlMessageType::RESIZE_CHAIN ) {
836+ // Multi-worker chain resize: apply our config and forward
837+ int my_idx = ctrl_msg.worker_index ;
838+ if (my_idx >= 0 && my_idx < ctrl_msg.total_workers ) {
839+ dist_config.split_layer = ctrl_msg.ranges [my_idx].start_layer ;
840+ dist_config.end_layer = ctrl_msg.ranges [my_idx].end_layer ;
841+ dist_config.is_tail = ctrl_msg.ranges [my_idx].is_tail ;
842+ start_layer = dist_config.split_layer ;
843+
844+ std::cout << " Worker " << my_idx << " : Layer config updated. Now processing layers "
845+ << start_layer << " to " << dist_config.end_layer - 1
846+ << (dist_config.is_tail ? " (tail)" : " " ) << std::endl;
847+ }
848+
849+ // Forward to next worker if there are more workers
850+ if (my_idx + 1 < ctrl_msg.total_workers ) {
851+ ctrl_msg.worker_index = my_idx + 1 ; // Increment for next worker
852+ dist_config.transport ->send_control (ctrl_msg);
853+ } else {
854+ // This is the last worker - send ACK back through the ring
855+ ControlMessage ack{};
856+ ack.type = ControlMessageType::ACK ;
857+ ack.split_layer = start_layer;
858+ ack.end_layer = dist_config.end_layer ;
859+ ack.is_tail = dist_config.is_tail ;
860+ dist_config.transport ->send_control (ack);
861+ }
862+ } else if (ctrl_msg.type == ControlMessageType::ACK ) {
863+ // Forward ACK to next node (back to master)
864+ dist_config.transport ->send_control (ctrl_msg);
865+ }
866+ }
867+
730868 // Receive header + data from Prev
731869 dist_config.transport ->recv_prev (transfer_buffer.data (), packet_size);
732870
871+ // Check if this is actually a control message (for inline control on kernel transport)
872+ uint16_t magic = 0 ;
873+ std::memcpy (&magic, transfer_buffer.data (), sizeof (magic));
874+ if (magic == CONTROL_MAGIC ) {
875+ ControlPacketHeader ctrl_pkt{};
876+ std::memcpy (&ctrl_pkt, transfer_buffer.data (), sizeof (ctrl_pkt));
877+ if (ctrl_pkt.msg .type == ControlMessageType::RESIZE_LAYERS ) {
878+ update_layer_config (ctrl_pkt.msg );
879+ start_layer = dist_config.split_layer ;
880+ std::cout << " Worker: Layer config updated (inline). Now processing layers "
881+ << start_layer << " to " << dist_config.end_layer - 1 << std::endl;
882+
883+ // Send ACK back
884+ ControlMessage ack{};
885+ ack.type = ControlMessageType::ACK ;
886+ ack.split_layer = start_layer;
887+ ack.end_layer = dist_config.end_layer ;
888+ ack.is_tail = dist_config.is_tail ;
889+ dist_config.transport ->send_control (ack);
890+ } else if (ctrl_pkt.msg .type == ControlMessageType::RESIZE_CHAIN ) {
891+ // Multi-worker chain resize (inline)
892+ int my_idx = ctrl_pkt.msg .worker_index ;
893+ if (my_idx >= 0 && my_idx < ctrl_pkt.msg .total_workers ) {
894+ dist_config.split_layer = ctrl_pkt.msg .ranges [my_idx].start_layer ;
895+ dist_config.end_layer = ctrl_pkt.msg .ranges [my_idx].end_layer ;
896+ dist_config.is_tail = ctrl_pkt.msg .ranges [my_idx].is_tail ;
897+ start_layer = dist_config.split_layer ;
898+
899+ std::cout << " Worker " << my_idx << " : Layer config updated (inline). Now processing layers "
900+ << start_layer << " to " << dist_config.end_layer - 1
901+ << (dist_config.is_tail ? " (tail)" : " " ) << std::endl;
902+ }
903+
904+ // Forward to next worker if there are more workers
905+ if (my_idx + 1 < ctrl_pkt.msg .total_workers ) {
906+ ControlMessage fwd = ctrl_pkt.msg ;
907+ fwd.worker_index = my_idx + 1 ;
908+ dist_config.transport ->send_control (fwd);
909+ } else {
910+ // This is the last worker - send ACK back through the ring
911+ ControlMessage ack{};
912+ ack.type = ControlMessageType::ACK ;
913+ ack.split_layer = start_layer;
914+ ack.end_layer = dist_config.end_layer ;
915+ ack.is_tail = dist_config.is_tail ;
916+ dist_config.transport ->send_control (ack);
917+ }
918+ } else if (ctrl_pkt.msg .type == ControlMessageType::ACK ) {
919+ // Forward ACK to next node (back to master)
920+ dist_config.transport ->send_control (ctrl_pkt.msg );
921+ }
922+ continue ; // Skip normal processing for control packets
923+ }
924+
733925 std::memcpy (&header, transfer_buffer.data (), sizeof (PacketHeader));
734926 std::memcpy (x, transfer_buffer.data () + sizeof (PacketHeader), dim * sizeof (float ));
735927
0 commit comments