Skip to content
40 changes: 38 additions & 2 deletions src/inference/FloatTransformer.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -814,12 +814,18 @@ namespace Inference {
std::cout << "Worker started. Processing layers " << start_layer << " to " << dist_config.end_layer - 1 <<
std::endl;

// Track old layer range for KV cache transfer
int old_start = start_layer;
int old_end = dist_config.end_layer;

while (true) {
try {
// Check for control messages (non-blocking)
ControlMessage ctrl_msg{};
if (dist_config.transport->recv_control_nonblocking(ctrl_msg)) {
if (ctrl_msg.type == ControlMessageType::RESIZE_LAYERS) {
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;
update_layer_config(ctrl_msg);
start_layer = dist_config.split_layer;
std::cout << "Worker: Layer config updated. Now processing layers "
Expand All @@ -833,7 +839,9 @@ namespace Inference {
ack.is_tail = dist_config.is_tail;
dist_config.transport->send_control(ack);
} else if (ctrl_msg.type == ControlMessageType::RESIZE_CHAIN) {
// Multi-worker chain resize: apply our config and forward
// Multi-worker chain resize: save old range, then apply
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;
int my_idx = ctrl_msg.worker_index;
if (my_idx >= 0 && my_idx < ctrl_msg.total_workers) {
dist_config.split_layer = ctrl_msg.ranges[my_idx].start_layer;
Expand Down Expand Up @@ -875,6 +883,8 @@ namespace Inference {
ControlPacketHeader ctrl_pkt{};
std::memcpy(&ctrl_pkt, transfer_buffer.data(), sizeof(ctrl_pkt));
if (ctrl_pkt.msg.type == ControlMessageType::RESIZE_LAYERS) {
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;
update_layer_config(ctrl_pkt.msg);
start_layer = dist_config.split_layer;
std::cout << "Worker: Layer config updated (inline). Now processing layers "
Expand All @@ -888,7 +898,9 @@ namespace Inference {
ack.is_tail = dist_config.is_tail;
dist_config.transport->send_control(ack);
} else if (ctrl_pkt.msg.type == ControlMessageType::RESIZE_CHAIN) {
// Multi-worker chain resize (inline)
// Multi-worker chain resize (inline): save old range first
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;
int my_idx = ctrl_pkt.msg.worker_index;
if (my_idx >= 0 && my_idx < ctrl_pkt.msg.total_workers) {
dist_config.split_layer = ctrl_pkt.msg.ranges[my_idx].start_layer;
Expand Down Expand Up @@ -922,6 +934,30 @@ namespace Inference {
continue; // Skip normal processing for control packets
}

// Check for KV cache transfer bundle
if (magic == KV_TRANSFER_MAGIC) {
KvTransferHeader kv_hdr{};
std::memcpy(&kv_hdr, transfer_buffer.data(), sizeof(kv_hdr));
const int kv_dim = config.dim * config.n_kv_heads / config.n_heads;

std::cout << "Worker: KV transfer received (" << kv_hdr.num_slices
<< " slices, pos=" << kv_hdr.pos << ")" << std::endl;

handle_kv_transfer(
state.key_cache.data(), state.value_cache.data(),
kv_dim, config.seq_len,
old_start, old_end,
dist_config.split_layer, dist_config.end_layer,
kv_hdr);

// Update old range to current (transfer complete)
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;

Comment on lines +942 to +956

Copilot AI Feb 9, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The KV-transfer path accepts kv_hdr and immediately calls handle_kv_transfer() without checking that kv_hdr.pos <= config.seq_len and kv_hdr.kv_dim matches the computed kv_dim. If either is wrong (corrupted packet / version skew), handle_kv_transfer() will read/write incorrect sizes and desynchronize the stream. Add a cheap validation here (and/or ensure handle_kv_transfer() robustly drains on invalid headers).

Suggested change
std::cout << "Worker: KV transfer received (" << kv_hdr.num_slices
<< " slices, pos=" << kv_hdr.pos << ")" << std::endl;
handle_kv_transfer(
state.key_cache.data(), state.value_cache.data(),
kv_dim, config.seq_len,
old_start, old_end,
dist_config.split_layer, dist_config.end_layer,
kv_hdr);
// Update old range to current (transfer complete)
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;
// Validate KV transfer header to avoid stream desynchronization on bad packets
if (kv_hdr.pos > config.seq_len || kv_hdr.kv_dim != kv_dim) {
std::cerr << "Worker: invalid KV transfer header received: "
<< "pos=" << kv_hdr.pos
<< " (max " << config.seq_len << "), "
<< "kv_dim=" << kv_hdr.kv_dim
<< " (expected " << kv_dim << ")" << std::endl;
// Drop this KV transfer packet and continue with the next one
continue;
}
std::cout << "Worker: KV transfer received (" << kv_hdr.num_slices
<< " slices, pos=" << kv_hdr.pos << ")" << std::endl;
handle_kv_transfer(
state.key_cache.data(), state.value_cache.data(),
kv_dim, config.seq_len,
old_start, old_end,
dist_config.split_layer, dist_config.end_layer,
kv_hdr);
// Update old range to current (transfer complete)
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;

Copilot uses AI. Check for mistakes.
std::cout << "Worker: KV transfer complete" << std::endl;
continue; // Skip normal processing for KV transfer packets
}

std::memcpy(&header, transfer_buffer.data(), sizeof(PacketHeader));
std::memcpy(x, transfer_buffer.data() + sizeof(PacketHeader), dim * sizeof(float));

Expand Down
3 changes: 3 additions & 0 deletions src/inference/FloatTransformer.h
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,9 @@ namespace Inference {

void clear_kv_cache() override;

float* get_key_cache() override { return state.key_cache.data(); }
float* get_value_cache() override { return state.value_cache.data(); }

private:
FloatTransformerWeights weights{};
FloatRunState state{};
Expand Down
40 changes: 38 additions & 2 deletions src/inference/QuantizedTransformer.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1080,12 +1080,18 @@ namespace Inference {
std::cout << "Worker started. Processing layers " << start_layer << " to " << dist_config.end_layer - 1 <<
std::endl;

// Track old layer range for KV cache transfer
int old_start = start_layer;
int old_end = dist_config.end_layer;

while (true) {
try {
// Check for control messages (non-blocking)
ControlMessage ctrl_msg{};
if (dist_config.transport->recv_control_nonblocking(ctrl_msg)) {
if (ctrl_msg.type == ControlMessageType::RESIZE_LAYERS) {
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;
update_layer_config(ctrl_msg);
start_layer = dist_config.split_layer;
std::cout << "Worker: Layer config updated. Now processing layers "
Expand All @@ -1099,7 +1105,9 @@ namespace Inference {
ack.is_tail = dist_config.is_tail;
dist_config.transport->send_control(ack);
} else if (ctrl_msg.type == ControlMessageType::RESIZE_CHAIN) {
// Multi-worker chain resize: apply our config and forward
// Multi-worker chain resize: save old range, then apply
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;
int my_idx = ctrl_msg.worker_index;
if (my_idx >= 0 && my_idx < ctrl_msg.total_workers) {
dist_config.split_layer = ctrl_msg.ranges[my_idx].start_layer;
Expand Down Expand Up @@ -1141,6 +1149,8 @@ namespace Inference {
ControlPacketHeader ctrl_pkt{};
std::memcpy(&ctrl_pkt, transfer_buffer.data(), sizeof(ctrl_pkt));
if (ctrl_pkt.msg.type == ControlMessageType::RESIZE_LAYERS) {
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;
update_layer_config(ctrl_pkt.msg);
start_layer = dist_config.split_layer;
std::cout << "Worker: Layer config updated (inline). Now processing layers "
Expand All @@ -1154,7 +1164,9 @@ namespace Inference {
ack.is_tail = dist_config.is_tail;
dist_config.transport->send_control(ack);
} else if (ctrl_pkt.msg.type == ControlMessageType::RESIZE_CHAIN) {
// Multi-worker chain resize (inline)
// Multi-worker chain resize (inline): save old range first
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;
int my_idx = ctrl_pkt.msg.worker_index;
if (my_idx >= 0 && my_idx < ctrl_pkt.msg.total_workers) {
dist_config.split_layer = ctrl_pkt.msg.ranges[my_idx].start_layer;
Expand Down Expand Up @@ -1188,6 +1200,30 @@ namespace Inference {
continue; // Skip normal processing for control packets
}

// Check for KV cache transfer bundle
if (magic == KV_TRANSFER_MAGIC) {
KvTransferHeader kv_hdr{};
std::memcpy(&kv_hdr, transfer_buffer.data(), sizeof(kv_hdr));
const int kv_dim = config.dim * config.n_kv_heads / config.n_heads;

std::cout << "Worker: KV transfer received (" << kv_hdr.num_slices
<< " slices, pos=" << kv_hdr.pos << ")" << std::endl;

handle_kv_transfer(
state.key_cache.data(), state.value_cache.data(),
kv_dim, config.seq_len,
old_start, old_end,
dist_config.split_layer, dist_config.end_layer,
kv_hdr);

// Update old range to current (transfer complete)
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;

Comment on lines +1208 to +1222

Copilot AI Feb 9, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The KV-transfer path accepts kv_hdr and immediately calls handle_kv_transfer() without checking that kv_hdr.pos <= config.seq_len and kv_hdr.kv_dim matches the computed kv_dim. If either is wrong (corrupted packet / version skew), handle_kv_transfer() will read/write incorrect sizes and desynchronize the stream. Add a cheap validation here (and/or ensure handle_kv_transfer() robustly drains on invalid headers).

Suggested change
std::cout << "Worker: KV transfer received (" << kv_hdr.num_slices
<< " slices, pos=" << kv_hdr.pos << ")" << std::endl;
handle_kv_transfer(
state.key_cache.data(), state.value_cache.data(),
kv_dim, config.seq_len,
old_start, old_end,
dist_config.split_layer, dist_config.end_layer,
kv_hdr);
// Update old range to current (transfer complete)
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;
// Validate header before using it to drive KV cache accesses
if (kv_hdr.pos > config.seq_len || kv_hdr.kv_dim != kv_dim) {
std::cerr << "Worker: invalid KV transfer header: "
<< "pos=" << kv_hdr.pos << " (max " << config.seq_len << "), "
<< "kv_dim=" << kv_hdr.kv_dim << " (expected " << kv_dim << ")"
<< std::endl;
// Skip this packet; do not modify KV cache or layer range
continue;
}
std::cout << "Worker: KV transfer received (" << kv_hdr.num_slices
<< " slices, pos=" << kv_hdr.pos << ")" << std::endl;
handle_kv_transfer(
state.key_cache.data(), state.value_cache.data(),
kv_dim, config.seq_len,
old_start, old_end,
dist_config.split_layer, dist_config.end_layer,
kv_hdr);
// Update old range to current (transfer complete)
old_start = dist_config.split_layer;
old_end = dist_config.end_layer;

Copilot uses AI. Check for mistakes.
std::cout << "Worker: KV transfer complete" << std::endl;
continue; // Skip normal processing for KV transfer packets
}

std::memcpy(&header, transfer_buffer.data(), sizeof(PacketHeader));
std::memcpy(x, transfer_buffer.data() + sizeof(PacketHeader), dim * sizeof(float));

Expand Down
3 changes: 3 additions & 0 deletions src/inference/QuantizedTransformer.h
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,9 @@ namespace Inference {

void clear_kv_cache() override;

float* get_key_cache() override { return state.key_cache.data(); }
float* get_value_cache() override { return state.value_cache.data(); }

private:
QuantizedTransformerWeights weights{};
QuantizedRunState state{};
Expand Down
157 changes: 157 additions & 0 deletions src/inference/Transformer.h
Original file line number Diff line number Diff line change
Expand Up @@ -54,9 +54,166 @@ namespace Inference {
dist_config.is_tail = msg.is_tail;
}

// Access KV cache data (for KV transfer protocol)
virtual float* get_key_cache() = 0;
virtual float* get_value_cache() = 0;

// Clear KV cache (needed after layer redistribution to avoid stale state)
virtual void clear_kv_cache() = 0;

// KV Cache Transfer: Master initiates the ring transfer after resize
// Sends departing layer KV data forward, then receives returning layers from tail
void initiate_kv_transfer(float *key_cache, float *value_cache,
int pos, int kv_dim, int seq_len,
int old_split, int new_split) {
if (!dist_config.transport || pos <= 0) return;

const size_t slice_data_size = static_cast<size_t>(pos) * kv_dim * sizeof(float);
const size_t packet_size = dist_config.transport->get_packet_size();

// Build bundle: layers master is giving up (old range that's no longer ours)
// Master owns [0, split_layer). Old: [0, old_split), New: [0, new_split)
std::vector<int32_t> departing_layers;
for (int l = new_split; l < old_split; l++) {
departing_layers.push_back(l); // Master lost these layers
}

// Send KV transfer header (padded to packet_size)
KvTransferHeader hdr{};
hdr.magic = KV_TRANSFER_MAGIC;
hdr.num_slices = static_cast<int32_t>(departing_layers.size());
hdr.pos = pos;
hdr.kv_dim = kv_dim;

std::vector<char> hdr_buf(packet_size, 0);
std::memcpy(hdr_buf.data(), &hdr, sizeof(hdr));
dist_config.transport->send_next(hdr_buf.data(), packet_size);

Copilot AI Feb 9, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

initiate_kv_transfer() assumes dist_config.transport->get_packet_size() is non-zero and large enough for KvTransferHeader. If packet_size is 0 (or < sizeof(KvTransferHeader)), hdr_buf becomes empty and the memcpy/send will be UB and will also desync receivers that are blocking on recv_prev(packet_size). Add a guard/fallback: ensure packet_size >= sizeof(KvTransferHeader) (e.g., compute the data packet size from config.dim when unset, or hard-fail with a clear error).

Copilot uses AI. Check for mistakes.

// Send each departing layer's KV data
for (int32_t layer_id : departing_layers) {
// Send layer_id
dist_config.transport->send_next(&layer_id, sizeof(layer_id));
// Send key cache for this layer
float *key_ptr = key_cache + static_cast<size_t>(layer_id) * seq_len * kv_dim;
dist_config.transport->send_next(key_ptr, slice_data_size);
// Send value cache for this layer
float *val_ptr = value_cache + static_cast<size_t>(layer_id) * seq_len * kv_dim;
dist_config.transport->send_next(val_ptr, slice_data_size);
}

// Now receive the return bundle (layers coming back from workers)
std::vector<char> recv_hdr_buf(packet_size);
dist_config.transport->recv_prev(recv_hdr_buf.data(), packet_size);

KvTransferHeader recv_hdr{};
std::memcpy(&recv_hdr, recv_hdr_buf.data(), sizeof(recv_hdr));

if (recv_hdr.magic == KV_TRANSFER_MAGIC && recv_hdr.num_slices > 0) {
for (int i = 0; i < recv_hdr.num_slices; i++) {
int32_t layer_id;
dist_config.transport->recv_prev(&layer_id, sizeof(layer_id));

// Check if this layer is in master's new range [0, new_split)
if (layer_id >= 0 && layer_id < new_split) {
float *key_ptr = key_cache + static_cast<size_t>(layer_id) * seq_len * kv_dim;
dist_config.transport->recv_prev(key_ptr, slice_data_size);
float *val_ptr = value_cache + static_cast<size_t>(layer_id) * seq_len * kv_dim;
dist_config.transport->recv_prev(val_ptr, slice_data_size);
} else {
// Discard (shouldn't happen after full ring rotation)
std::vector<char> discard(slice_data_size);
dist_config.transport->recv_prev(discard.data(), slice_data_size);
dist_config.transport->recv_prev(discard.data(), slice_data_size);
}

Copilot AI Feb 9, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The discard path allocates std::vector<char> discard(slice_data_size) inside the loop. For large pos * kv_dim, this can repeatedly allocate tens/hundreds of MB and drastically slow or OOM the process. Reuse a single scratch buffer outside the loop (or stream-discard into transfer_buffer) instead of allocating per slice.

Copilot uses AI. Check for mistakes.
}
}
}

// KV Cache Transfer: Worker handles an incoming KV bundle
// Extracts layers it needs, adds its departing layers, forwards the rest
void handle_kv_transfer(float *key_cache, float *value_cache,
int kv_dim, int seq_len,
int old_start, int old_end,
int new_start, int new_end,
const KvTransferHeader &incoming_hdr) {
const int pos = incoming_hdr.pos;
const size_t slice_data_size = static_cast<size_t>(pos) * kv_dim * sizeof(float);
const size_t packet_size = dist_config.transport->get_packet_size();

// Receive all incoming slices
struct KvSlice {
int32_t layer_id;
std::vector<float> key_data;
std::vector<float> value_data;
};
std::vector<KvSlice> incoming_slices;

for (int i = 0; i < incoming_hdr.num_slices; i++) {
KvSlice slice;
dist_config.transport->recv_prev(&slice.layer_id, sizeof(slice.layer_id));
slice.key_data.resize(static_cast<size_t>(pos) * kv_dim);
slice.value_data.resize(static_cast<size_t>(pos) * kv_dim);
dist_config.transport->recv_prev(slice.key_data.data(), slice_data_size);

Copilot AI Feb 9, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

handle_kv_transfer() trusts incoming_hdr.pos/incoming_hdr.kv_dim/incoming_hdr.num_slices without validating them against seq_len, kv_dim, and reasonable bounds. If these values are corrupted or a node runs a mismatched version, this can cause out-of-bounds writes into the KV caches or attempt absurd allocations. Add validation (and if invalid, drain/forward the corresponding bytes to keep the stream aligned before returning/erroring).

Copilot uses AI. Check for mistakes.
dist_config.transport->recv_prev(slice.value_data.data(), slice_data_size);
incoming_slices.push_back(std::move(slice));
}

Copilot AI Feb 9, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

handle_kv_transfer() buffers all incoming KV slices into incoming_slices (each slice holds 2x pos*kv_dim floats). With realistic seq lengths / kv dims, this can easily explode memory (GBs) during a resize and crash workers. Consider changing the transfer protocol to allow streaming (e.g., bundle byte length or end-of-bundle sentinel) so workers can forward slices without materializing the entire bundle in RAM.

Copilot uses AI. Check for mistakes.

// Extract layers this worker needs (in new range but NOT in old range)
for (auto &slice : incoming_slices) {
if (slice.layer_id >= new_start && slice.layer_id < new_end) {
// Copy into our KV cache
float *key_ptr = key_cache + static_cast<size_t>(slice.layer_id) * seq_len * kv_dim;
float *val_ptr = value_cache + static_cast<size_t>(slice.layer_id) * seq_len * kv_dim;
std::memcpy(key_ptr, slice.key_data.data(), slice_data_size);
std::memcpy(val_ptr, slice.value_data.data(), slice_data_size);
slice.layer_id = -1; // Mark as consumed

Copilot AI Feb 9, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The comment says the worker extracts layers it needs “in new range but NOT in old range”, but the condition only checks layer_id >= new_start && layer_id < new_end (no old-range exclusion). Either update the condition to match the intended behavior or adjust the comment so future readers don’t assume an invariant that isn’t enforced.

Copilot uses AI. Check for mistakes.
}
}

// Build forward bundle: unconsumed incoming slices + our departing layers
std::vector<KvSlice> forward_slices;

// Add unconsumed incoming slices
for (auto &slice : incoming_slices) {
if (slice.layer_id >= 0) {
forward_slices.push_back(std::move(slice));
}
}

// Add our departing layers (in old range but NOT in new range)
for (int l = old_start; l < old_end; l++) {
if (l < new_start || l >= new_end) {
KvSlice slice;
slice.layer_id = l;
slice.key_data.resize(static_cast<size_t>(pos) * kv_dim);
slice.value_data.resize(static_cast<size_t>(pos) * kv_dim);
float *key_ptr = key_cache + static_cast<size_t>(l) * seq_len * kv_dim;
float *val_ptr = value_cache + static_cast<size_t>(l) * seq_len * kv_dim;
std::memcpy(slice.key_data.data(), key_ptr, slice_data_size);
std::memcpy(slice.value_data.data(), val_ptr, slice_data_size);
forward_slices.push_back(std::move(slice));
}
}

// Send forward bundle header
KvTransferHeader fwd_hdr{};
fwd_hdr.magic = KV_TRANSFER_MAGIC;
fwd_hdr.num_slices = static_cast<int32_t>(forward_slices.size());
fwd_hdr.pos = pos;
fwd_hdr.kv_dim = kv_dim;

std::vector<char> hdr_buf(packet_size, 0);
std::memcpy(hdr_buf.data(), &fwd_hdr, sizeof(fwd_hdr));
dist_config.transport->send_next(hdr_buf.data(), packet_size);

// Send each forwarded slice
for (auto &slice : forward_slices) {
dist_config.transport->send_next(&slice.layer_id, sizeof(slice.layer_id));
dist_config.transport->send_next(slice.key_data.data(), slice_data_size);
dist_config.transport->send_next(slice.value_data.data(), slice_data_size);
}
}

// Factory method to create the appropriate Transformer (Float or Quantized) based on file
static std::unique_ptr<Transformer> create(const std::string &checkpoint_path);

Expand Down
14 changes: 14 additions & 0 deletions src/inference/Transport.h
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,20 @@ namespace Inference {
ControlMessage msg;
} __attribute__((packed));

// KV Cache Transfer Protocol
constexpr uint16_t KV_TRANSFER_MAGIC = 0xCAFE;

// Maximum layers in a single model
constexpr int MAX_LAYERS = 128;

// Header for KV cache transfer bundles (packed for network transmission)
struct KvTransferHeader {
uint16_t magic; // KV_TRANSFER_MAGIC
int32_t num_slices; // number of layer slices in this bundle
int32_t pos; // filled sequence positions (how much cache is valid)
int32_t kv_dim; // KV dimension per layer
} __attribute__((packed));

class Transport {
protected:
size_t packet_size_ = 0; // Data packet size (header + dim * sizeof(float))
Expand Down
Loading