-
Notifications
You must be signed in to change notification settings - Fork 0
feat(inference): implement KV cache transfer protocol for distributed mode #8
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from 1 commit
a11d0a2
881380e
d7742dd
b1a7f73
5007f07
1cf25b6
6c4e528
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
|
|
@@ -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 " | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
@@ -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; | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
@@ -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 " | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
@@ -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; | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
@@ -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
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| 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; |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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); | ||
|
||
|
|
||
| // 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); | ||
| } | ||
|
||
| } | ||
| } | ||
| } | ||
|
|
||
| // 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); | ||
|
||
| dist_config.transport->recv_prev(slice.value_data.data(), slice_data_size); | ||
| incoming_slices.push_back(std::move(slice)); | ||
| } | ||
|
||
|
|
||
| // 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 | ||
|
||
| } | ||
| } | ||
|
|
||
| // 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); | ||
|
|
||
|
|
||
There was a problem hiding this comment.
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_hdrand immediately callshandle_kv_transfer()without checking thatkv_hdr.pos <= config.seq_lenandkv_hdr.kv_dimmatches the computedkv_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 ensurehandle_kv_transfer()robustly drains on invalid headers).