Skip to content

Commit d96c8bd

Browse files
authored
Merge pull request #7 from Harikeshav-R/dynamic-layer
Add the ability to dynamically resize layer processing between inference calls for load balancing purposes
2 parents 5aa5ad1 + d0fcacc commit d96c8bd

14 files changed

Lines changed: 856 additions & 39 deletions

src/inference/FloatTransformer.cpp

Lines changed: 212 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -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

src/inference/FloatTransformer.h

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,8 @@ namespace Inference {
4747

4848
void worker_loop() override;
4949

50+
void clear_kv_cache() override;
51+
5052
private:
5153
FloatTransformerWeights weights{};
5254
FloatRunState state{};

src/inference/KernelTransport.cpp

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -162,4 +162,29 @@ namespace Inference {
162162
// Recv from prev node - Update prev_ip
163163
recv_internal(data, size, true);
164164
}
165+
166+
void KernelTransport::send_control(const ControlMessage &msg) {
167+
if (next_ip.empty()) throw std::runtime_error("Next IP not configured for KernelTransport control");
168+
169+
// For kernel transport, control messages use the same mechanism as regular data
170+
// Pad to packet_size_ for consistent recv_prev handling
171+
ControlPacketHeader pkt{CONTROL_MAGIC, msg};
172+
173+
if (packet_size_ > 0 && packet_size_ >= sizeof(pkt)) {
174+
std::vector<char> buffer(packet_size_, 0);
175+
std::memcpy(buffer.data(), &pkt, sizeof(pkt));
176+
send_next(buffer.data(), buffer.size());
177+
} else {
178+
send_next(&pkt, sizeof(pkt));
179+
}
180+
}
181+
182+
bool KernelTransport::recv_control_nonblocking(ControlMessage &msg) {
183+
// Kernel transport uses blocking ioctl for data receive
184+
// For control messages, we can't easily do non-blocking check without modifying kernel module
185+
// For now, control messages for kernel transport are checked inline in worker_loop
186+
// by inspecting the first bytes of received data before processing
187+
(void)msg;
188+
return false; // Non-blocking control not directly supported; handled at higher level
189+
}
165190
}

src/inference/KernelTransport.h

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,10 @@ namespace Inference {
2828

2929
void send_multipart_next(const void *header, size_t header_size, const void *data, size_t data_size) override;
3030

31+
// Control channel for dynamic layer resizing
32+
void send_control(const ControlMessage &msg) override;
33+
bool recv_control_nonblocking(ControlMessage &msg) override;
34+
3135
private:
3236
void set_destination(const std::string &ip, int target_port) const;
3337

0 commit comments

Comments
 (0)