diff --git a/tpu_sync/core/reshard_receive_session.cc b/tpu_sync/core/reshard_receive_session.cc index 8b404a3d..d1fc796c 100644 --- a/tpu_sync/core/reshard_receive_session.cc +++ b/tpu_sync/core/reshard_receive_session.cc @@ -146,7 +146,7 @@ absl::Status ReshardReceiveSession::ValidatePlan( absl::Status ReshardReceiveSession::ValidateReceiverCoverage( const kv_cache::KVCacheManagerBase& base, const ::tpu_sync::rpc::StartTransferRequest& plan) { - constexpr int64_t kMaxExpandedRepeats = 1 << 20; + constexpr int64_t kMaxExpandedRepeats = 1 << 24; struct GroupView { std::vector pool_indices; diff --git a/tpu_sync/kv_cache/reshard/pool_reshard_planner.cc b/tpu_sync/kv_cache/reshard/pool_reshard_planner.cc index 677b48b2..9506164c 100644 --- a/tpu_sync/kv_cache/reshard/pool_reshard_planner.cc +++ b/tpu_sync/kv_cache/reshard/pool_reshard_planner.cc @@ -128,8 +128,445 @@ struct TagPrecheck { int64_t dst_live = 0; std::vector src_segments; std::vector dst_segments; + // Both live maps are one region at physical offset zero, so logical and + // physical block offsets coincide and a strided span needs no translation. + bool identity_maps = false; }; +// True when a pool's live bytes are one contiguous region at physical +// offset zero. +bool IdentityLiveMap(const std::vector& segments, + int64_t live) { + return segments.size() == 1 && segments[0].logical_offset == 0 && + segments[0].physical_offset == 0 && segments[0].size == live; +} + +// One source unit's byte-span declaration for the tag being planned. +using DeclaredSpans = + std::vector>; +using PeersByUnit = + absl::btree_map; +using SchedulesByUnit = std::map, + RequestBlockRegistry::RaidenIdLess>; +// (destination peer, source block, destination block) pushes, by sender and +// then by destination unit. +using TransferPairs = std::map< + RaidenId, + std::map>, + RequestBlockRegistry::RaidenIdLess>, + RequestBlockRegistry::RaidenIdLess>; + +// Re-bases one strided request-global span (one head slice repeated per +// token) to its destination page. The span must lie inside one destination +// page, so a page-aligned clip drops it whole or keeps it whole. +absl::Status RebaseStridedGlobalSpan(const PoolByteSpan& span, + const RaidenId& unit, int64_t tag_skip, + int64_t dst_live, int64_t* tag_global_end, + int64_t* tag_clipped_bytes, + std::vector* out) { + if (span.src_stride_bytes <= 0 || span.dst_stride_bytes <= 0) { + return absl::InvalidArgumentError( + absl::StrCat("Strided global-space byte spans require positive strides " + "(declared by ", + PythonRepr(unit), ")")); + } + const int64_t in_page = span.dst_offset_bytes % dst_live; + const int64_t extent = + (span.count - 1) * span.dst_stride_bytes + span.size_bytes; + if (in_page + extent > dst_live) { + return absl::InvalidArgumentError( + absl::StrCat("Strided global-space byte spans must lie inside one " + "destination page (declared by ", + PythonRepr(unit), ")")); + } + *tag_global_end = std::max(*tag_global_end, span.dst_offset_bytes + extent); + if (span.dst_offset_bytes < tag_skip) { + *tag_clipped_bytes += span.size_bytes * span.count; + return absl::OkStatus(); + } + PoolByteSpan strided_span = span; + strided_span.dst_block_index = (span.dst_offset_bytes - tag_skip) / dst_live; + strided_span.dst_offset_bytes = in_page; + out->push_back(strided_span); + return absl::OkStatus(); +} + +// Splits one plain request-global span at destination page boundaries, after +// the page-aligned `tag_skip` clip: a span wholly below the cut is dropped +// (its source bytes are never read), a straddling span is trimmed, and +// survivors re-base into the clipped space. +void SplitPlainGlobalSpan(const PoolByteSpan& span, int64_t tag_skip, + int64_t dst_live, int64_t* tag_global_end, + int64_t* tag_clipped_bytes, + std::vector* out) { + *tag_global_end = + std::max(*tag_global_end, span.dst_offset_bytes + span.size_bytes); + int64_t clip_advance = 0; + if (span.dst_offset_bytes + span.size_bytes <= tag_skip) { + *tag_clipped_bytes += span.size_bytes; + return; + } + if (span.dst_offset_bytes < tag_skip) { + clip_advance = tag_skip - span.dst_offset_bytes; + *tag_clipped_bytes += clip_advance; + } + int64_t remaining = span.size_bytes - clip_advance; + int64_t global_offset = span.dst_offset_bytes + clip_advance - tag_skip; + int64_t src_offset = span.src_offset_bytes + clip_advance; + while (remaining > 0) { + const int64_t page_index = global_offset / dst_live; + const int64_t in_page = global_offset % dst_live; + const int64_t take = std::min(remaining, dst_live - in_page); + PoolByteSpan split_span; + split_span.src_block_ordinal = span.src_block_ordinal; + split_span.src_offset_bytes = src_offset; + split_span.dst_block_index = page_index; + split_span.dst_offset_bytes = in_page; + split_span.size_bytes = take; + split_span.dst_unit_ordinal = span.dst_unit_ordinal; + out->push_back(split_span); + global_offset += take; + src_offset += take; + remaining -= take; + } +} + +// T3.4: destination-page-agnostic declarations (dst_space_version=1) +// split at destination page boundaries here, where the destination +// geometry is known. A page-aligned dst_skip_bytes clip applies in the +// same pass: spans wholly below the cut are dropped (their source bytes +// are never read), the straddling span is trimmed, and survivors +// re-base into the clipped space — skip is a dst_live multiple, so +// in-page offsets are unchanged and only page indices shift. +// `declared` then points at the rewritten registrations, which `storage` +// owns. +absl::Status SplitGlobalSpaceSpans(const std::string& plan_tag, + int64_t tag_skip, int64_t dst_live, + DeclaredSpans* declared, + std::vector* storage, + int64_t* tag_clipped_bytes) { + storage->reserve(declared->size()); + DeclaredSpans converted; + int64_t tag_global_end = 0; + for (const auto& [unit, entry] : *declared) { + if (entry->dst_space_version == 0) { + if (tag_skip > 0) { + return absl::InvalidArgumentError(absl::StrCat( + "dst_skip_bytes requires destination-page-agnostic " + "(dst_space_version=1) declarations for tag ", + PyStrRepr(plan_tag), " (declared by ", PythonRepr(unit), ")")); + } + converted.emplace_back(unit, entry); + continue; + } + PoolSpanRegistration split_entry = *entry; + split_entry.spans.clear(); + split_entry.dst_space_version = 0; + for (const PoolByteSpan& span : entry->spans) { + if (span.dst_block_index != 0) { + return absl::InvalidArgumentError(absl::StrCat( + "Global-space byte spans must leave dst_block_index zero " + "(declared by ", + PythonRepr(unit), ")")); + } + if (span.count > 1) { + absl::Status status = RebaseStridedGlobalSpan( + span, unit, tag_skip, dst_live, &tag_global_end, tag_clipped_bytes, + &split_entry.spans); + if (!status.ok()) return status; + continue; + } + if (span.src_stride_bytes != 0 || span.dst_stride_bytes != 0) { + return absl::InvalidArgumentError(absl::StrCat( + "Global-space byte spans must be plain contiguous ranges " + "(declared by ", + PythonRepr(unit), ")")); + } + SplitPlainGlobalSpan(span, tag_skip, dst_live, &tag_global_end, + tag_clipped_bytes, &split_entry.spans); + } + storage->push_back(std::move(split_entry)); + converted.emplace_back(unit, &storage->back()); + } + if (tag_skip > 0 && tag_skip >= tag_global_end) { + return absl::InvalidArgumentError(absl::StrCat( + "dst_skip_bytes removes the entire tag ", PyStrRepr(plan_tag), + ": skip=", tag_skip, ", declared_extent=", tag_global_end, + "; a full local hit must not reach the planner")); + } + *declared = std::move(converted); + return absl::OkStatus(); +} + +// Validate: expand every span's uniform repeats; for every destination +// unit, the union of the repeats routed to it must cover each +// destination block exactly once as a prefix-shaped extent. A span +// without a dst_unit_ordinal (absent on the wire) is replicated to every +// destination; a span with an ordinal belongs to that destination only +// (sharded destinations, e.g. head-split KV caches). Declared/covered +// byte accounting counts each span repeat once however many +// destinations it reaches. +// On success `extents` holds the covered byte prefix of each destination +// block. +absl::Status ValidateSpanCoverage(const DeclaredSpans& declared, + const std::string& plan_tag, size_t num_dst, + size_t num_dst_blocks, int64_t src_live, + int64_t dst_live, int64_t tag_clipped_bytes, + std::vector* extents) { + std::vector>>> + coverage_by_dst(num_dst, + std::vector>>( + num_dst_blocks)); + int64_t expanded_repeats = 0; + int64_t declared_total = 0; + int64_t covered_total = 0; + int64_t replicated_total = 0; + for (const auto& [unit, entry] : declared) { + declared_total += entry->declared_bytes; + for (const PoolByteSpan& span : entry->spans) { + if (span.dst_unit_ordinal >= static_cast(num_dst)) { + return absl::InvalidArgumentError(absl::StrCat( + "Byte span dst_unit_ordinal ", span.dst_unit_ordinal, + " exceeds the transfer's ", num_dst, " destination units for tag ", + PyStrRepr(plan_tag), " (declared by ", PythonRepr(unit), ")")); + } + if (span.dst_block_index >= static_cast(num_dst_blocks)) { + return absl::InvalidArgumentError( + absl::StrCat("Byte span destination index ", span.dst_block_index, + " exceeds the transfer's ", num_dst_blocks, + " destination blocks for tag ", PyStrRepr(plan_tag), + " (declared by ", PythonRepr(unit), ")")); + } + const int64_t src_end = span.src_offset_bytes + + (span.count - 1) * span.src_stride_bytes + + span.size_bytes; + if (src_end > src_live) { + return absl::InvalidArgumentError(absl::StrCat( + "Byte span exceeds its source block live bytes: end=", src_end, + ", live=", src_live, " (declared by ", PythonRepr(unit), ")")); + } + const int64_t dst_end = span.dst_offset_bytes + + (span.count - 1) * span.dst_stride_bytes + + span.size_bytes; + if (dst_end > dst_live) { + return absl::InvalidArgumentError(absl::StrCat( + "Byte span exceeds its destination block live bytes: end=", dst_end, + ", live=", dst_live, " (declared by ", PythonRepr(unit), ")")); + } + expanded_repeats += span.count; + if (expanded_repeats > kMaxLiveSegments) { + return absl::InvalidArgumentError( + "Byte span plan exceeds the repeat expansion bound"); + } + for (int32_t repeat = 0; repeat < span.count; ++repeat) { + const int64_t start = + span.dst_offset_bytes + repeat * span.dst_stride_bytes; + if (span.dst_unit_ordinal < 0) { + for (size_t d = 0; d < num_dst; ++d) { + coverage_by_dst[d][span.dst_block_index].emplace_back( + start, start + span.size_bytes); + } + replicated_total += span.size_bytes; + } else { + coverage_by_dst[span.dst_unit_ordinal][span.dst_block_index] + .emplace_back(start, start + span.size_bytes); + } + covered_total += span.size_bytes; + } + } + } + for (size_t d = 0; d < num_dst; ++d) { + // Single-destination plans keep the historical error strings; the + // destination suffix appears only for multi-destination plans. + const std::string dst_suffix = + num_dst > 1 ? absl::StrCat(" at destination unit ", d) : ""; + std::vector dst_extents; + for (size_t index = 0; index < num_dst_blocks; ++index) { + std::vector>& intervals = + coverage_by_dst[d][index]; + std::sort(intervals.begin(), intervals.end()); + if (intervals.empty()) { + return absl::InvalidArgumentError( + absl::StrCat("Destination block index ", index, + " has no declared coverage for tag ", + PyStrRepr(plan_tag), dst_suffix)); + } + int64_t covered_until = 0; + for (const auto& [start, end] : intervals) { + if (start != covered_until) { + const bool overlap = start < covered_until; + return absl::InvalidArgumentError( + absl::StrCat("Declared byte spans have a destination coverage ", + overlap ? "overlap" : "gap", " at byte ", + std::min(start, covered_until), + " of destination block index ", index, dst_suffix)); + } + covered_until = end; + } + if (index != num_dst_blocks - 1 && covered_until != dst_live) { + return absl::InvalidArgumentError(absl::StrCat( + "Destination block index ", index, " is covered to ", covered_until, + " of ", dst_live, + " live bytes; only the final block may be partial", dst_suffix)); + } + dst_extents.push_back(covered_until); + } + if (d == 0) { + *extents = std::move(dst_extents); + } else if (dst_extents != *extents) { + // One extent vector per group travels to every receiver; sharded + // destinations must therefore cover the same byte prefix of every + // destination block (true for head-split caches by construction). + return absl::InvalidArgumentError(absl::StrCat( + "Destination units must share uniform coverage extents for tag ", + PyStrRepr(plan_tag), "; destination unit ", d, + " differs from destination unit 0")); + } + } + if (declared_total != covered_total + tag_clipped_bytes) { + if (tag_clipped_bytes == 0) { + return absl::InvalidArgumentError( + absl::StrCat("Declared byte totals disagree with coverage: declared=", + declared_total, ", covered=", covered_total)); + } + return absl::InvalidArgumentError( + absl::StrCat("Declared byte totals disagree with coverage: declared=", + declared_total, ", covered=", covered_total, + ", clipped=", tag_clipped_bytes)); + } + { + // Every destination covers extent_sum bytes; replicated spans are + // counted once in covered_total but reach all num_dst destinations. + int64_t extent_sum = 0; + for (int64_t extent : *extents) extent_sum += extent; + if (static_cast(num_dst) * extent_sum != + covered_total + static_cast(num_dst - 1) * replicated_total) { + return absl::InternalError("byte coverage accounting failed"); + } + } + return absl::OkStatus(); +} + +// Records the block pairs `src_unit` pushes to each of `targets`. +void AddTransferPairs(const RaidenId& src_unit, + const std::vector& targets, + const PeersByUnit& dst_peers, int64_t src_block_id, + int64_t dst_block_id, + TransferPairs* transfer_pairs_per_sender) { + auto& sender_pairs = (*transfer_pairs_per_sender)[src_unit]; + for (const RaidenId& dst_unit_id : targets) { + sender_pairs[dst_unit_id].insert( + std::make_tuple(dst_peers.at(dst_unit_id), src_block_id, dst_block_id)); + } +} + +// Bind ordinals and destination indices to physical ids and emit. +absl::Status EmitSpanSchedule( + const DeclaredSpans& declared, const std::vector& dst_units, + const PeersByUnit& dst_peers, const std::vector& dst_ids, + const TagPrecheck& precheck, int32_t pool_group, SchedulesByUnit* schedules, + int64_t* emitted_chunks, TransferPairs* transfer_pairs_per_sender) { + struct OrderedSpan { + const PoolByteSpan* span; + RaidenId unit; + const PoolSpanRegistration* entry; + }; + std::vector ordered_spans; + for (const auto& [unit, entry] : declared) { + for (const PoolByteSpan& span : entry->spans) { + ordered_spans.push_back(OrderedSpan{&span, unit, entry}); + } + } + std::stable_sort( + ordered_spans.begin(), ordered_spans.end(), + [](const OrderedSpan& a, const OrderedSpan& b) { + return std::tie(a.span->dst_block_index, a.span->dst_offset_bytes) < + std::tie(b.span->dst_block_index, b.span->dst_offset_bytes); + }); + for (const OrderedSpan& ordered : ordered_spans) { + const PoolByteSpan& span = *ordered.span; + const RaidenId& src_unit = ordered.unit; + const int64_t src_block_id = + ordered.entry->block_ids[span.src_block_ordinal]; + const int64_t dst_block_id = dst_ids[span.dst_block_index]; + // A span without an ordinal replicates its chunks to every + // destination (entries differ only in dst_peer); a span with an + // ordinal reaches that destination only. + std::vector targets; + if (span.dst_unit_ordinal < 0) { + targets = dst_units; + } else { + targets.push_back(dst_units[span.dst_unit_ordinal]); + } + if (span.count > 1 && precheck.identity_maps) { + // Logical and physical offsets coincide, so a strided span (one head + // slice per token) travels as one strided entry per destination + // instead of one entry per token. Its bounds were checked above. + *emitted_chunks += static_cast(targets.size()); + if (*emitted_chunks > kMaxLiveSegments) { + return absl::InvalidArgumentError( + "Byte-span plan exceeds the live-region expansion bound"); + } + for (const RaidenId& dst_unit_id : targets) { + ScheduleEntry schedule_entry; + schedule_entry.dst_peer = dst_peers.at(dst_unit_id); + schedule_entry.dst_shard_idx = 0; + schedule_entry.dst_offset_bytes = span.dst_offset_bytes; + schedule_entry.src_offset_bytes = span.src_offset_bytes; + schedule_entry.size_bytes = span.size_bytes; + schedule_entry.src_block_id = src_block_id; + schedule_entry.dst_block_id = dst_block_id; + schedule_entry.src_stride_bytes = span.src_stride_bytes; + schedule_entry.dst_stride_bytes = span.dst_stride_bytes; + schedule_entry.count = span.count; + schedule_entry.layer_idx = 0; + schedule_entry.pool_group = pool_group; + (*schedules)[src_unit].push_back(std::move(schedule_entry)); + } + AddTransferPairs(src_unit, targets, dst_peers, src_block_id, dst_block_id, + transfer_pairs_per_sender); + continue; + } + for (int32_t repeat = 0; repeat < span.count; ++repeat) { + const int64_t src_offset = + span.src_offset_bytes + repeat * span.src_stride_bytes; + const int64_t dst_offset = + span.dst_offset_bytes + repeat * span.dst_stride_bytes; + auto translated = + TranslateLiveCopy(precheck.src_segments, precheck.dst_segments, + src_offset, dst_offset, span.size_bytes); + if (!translated.ok()) return translated.status(); + *emitted_chunks += static_cast(translated->size()) * + static_cast(targets.size()); + if (*emitted_chunks > kMaxLiveSegments) { + return absl::InvalidArgumentError( + "Byte-span plan exceeds the live-region expansion bound"); + } + for (const LiveCopyChunk& chunk : *translated) { + for (const RaidenId& dst_unit_id : targets) { + ScheduleEntry schedule_entry; + schedule_entry.dst_peer = dst_peers.at(dst_unit_id); + schedule_entry.dst_shard_idx = 0; + schedule_entry.dst_offset_bytes = chunk.dst_physical; + schedule_entry.src_offset_bytes = chunk.src_physical; + schedule_entry.size_bytes = chunk.size; + schedule_entry.src_block_id = src_block_id; + schedule_entry.dst_block_id = dst_block_id; + schedule_entry.src_stride_bytes = 0; + schedule_entry.dst_stride_bytes = 0; + schedule_entry.count = 1; + schedule_entry.layer_idx = 0; + schedule_entry.pool_group = pool_group; + (*schedules)[src_unit].push_back(std::move(schedule_entry)); + } + } + } + AddTransferPairs(src_unit, targets, dst_peers, src_block_id, dst_block_id, + transfer_pairs_per_sender); + } + return absl::OkStatus(); +} + } // namespace absl::StatusOr BuildPoolReshardPlan( @@ -552,6 +989,9 @@ absl::StatusOr BuildPoolReshardPlan( } precheck.src_segments = std::move(src_segment_maps[0]); precheck.dst_segments = std::move(dst_segment_maps[0]); + precheck.identity_maps = + IdentityLiveMap(precheck.src_segments, precheck.src_live) && + IdentityLiveMap(precheck.dst_segments, precheck.dst_live); tag_precheck.push_back(std::move(precheck)); } @@ -618,88 +1058,12 @@ absl::StatusOr BuildPoolReshardPlan( plan_tag, ", req_id=", req_id, ", uuid=", uuid)); } - // T3.4: destination-page-agnostic declarations (dst_space_version=1) - // split at destination page boundaries here, where the destination - // geometry is known. A page-aligned dst_skip_bytes clip applies in the - // same pass: spans wholly below the cut are dropped (their source bytes - // are never read), the straddling span is trimmed, and survivors - // re-base into the clipped space — skip is a dst_live multiple, so - // in-page offsets are unchanged and only page indices shift. std::vector converted_storage; - converted_storage.reserve(declared.size()); - std::vector> converted; int64_t tag_clipped_bytes = 0; - int64_t tag_global_end = 0; - for (const auto& [unit, entry] : declared) { - if (entry->dst_space_version == 0) { - if (tag_skip > 0) { - return absl::InvalidArgumentError(absl::StrCat( - "dst_skip_bytes requires destination-page-agnostic " - "(dst_space_version=1) declarations for tag ", - PyStrRepr(plan_tag), " (declared by ", PythonRepr(unit), ")")); - } - converted.emplace_back(unit, entry); - continue; - } - PoolSpanRegistration split_entry = *entry; - split_entry.spans.clear(); - split_entry.dst_space_version = 0; - for (const PoolByteSpan& span : entry->spans) { - if (span.count > 1 || span.src_stride_bytes != 0 || - span.dst_stride_bytes != 0) { - return absl::InvalidArgumentError(absl::StrCat( - "Global-space byte spans must be plain contiguous ranges " - "(declared by ", - PythonRepr(unit), ")")); - } - if (span.dst_block_index != 0) { - return absl::InvalidArgumentError(absl::StrCat( - "Global-space byte spans must leave dst_block_index zero " - "(declared by ", - PythonRepr(unit), ")")); - } - tag_global_end = - std::max(tag_global_end, span.dst_offset_bytes + span.size_bytes); - int64_t clip_advance = 0; - if (span.dst_offset_bytes + span.size_bytes <= tag_skip) { - tag_clipped_bytes += span.size_bytes; - continue; - } - if (span.dst_offset_bytes < tag_skip) { - clip_advance = tag_skip - span.dst_offset_bytes; - tag_clipped_bytes += clip_advance; - } - int64_t remaining = span.size_bytes - clip_advance; - int64_t global_offset = - span.dst_offset_bytes + clip_advance - tag_skip; - int64_t src_offset = span.src_offset_bytes + clip_advance; - while (remaining > 0) { - const int64_t page_index = global_offset / dst_live; - const int64_t in_page = global_offset % dst_live; - const int64_t take = std::min(remaining, dst_live - in_page); - PoolByteSpan split_span; - split_span.src_block_ordinal = span.src_block_ordinal; - split_span.src_offset_bytes = src_offset; - split_span.dst_block_index = page_index; - split_span.dst_offset_bytes = in_page; - split_span.size_bytes = take; - split_span.dst_unit_ordinal = span.dst_unit_ordinal; - split_entry.spans.push_back(split_span); - global_offset += take; - src_offset += take; - remaining -= take; - } - } - converted_storage.push_back(std::move(split_entry)); - converted.emplace_back(unit, &converted_storage.back()); - } - if (tag_skip > 0 && tag_skip >= tag_global_end) { - return absl::InvalidArgumentError(absl::StrCat( - "dst_skip_bytes removes the entire tag ", PyStrRepr(plan_tag), - ": skip=", tag_skip, ", declared_extent=", tag_global_end, - "; a full local hit must not reach the planner")); - } - declared = std::move(converted); + absl::Status split_status = + SplitGlobalSpaceSpans(plan_tag, tag_skip, dst_live, &declared, + &converted_storage, &tag_clipped_bytes); + if (!split_status.ok()) return split_status; for (size_t k = 0; k < precheck.selected.size(); ++k) { const int32_t pool_idx = precheck.selected[k]; @@ -724,230 +1088,19 @@ absl::StatusOr BuildPoolReshardPlan( } } - // Validate: expand every span's uniform repeats; for every destination - // unit, the union of the repeats routed to it must cover each - // destination block exactly once as a prefix-shaped extent. A span - // without a dst_unit_ordinal (absent on the wire) is replicated to every - // destination; a span with an ordinal belongs to that destination only - // (sharded destinations, e.g. head-split KV caches). Declared/covered - // byte accounting counts each span repeat once however many - // destinations it reaches. const size_t num_dst = request.dst_units.size(); - std::vector>>> - coverage_by_dst(num_dst, - std::vector>>( - dst_ids_g.size())); - int64_t expanded_repeats = 0; - int64_t declared_total = 0; - int64_t covered_total = 0; - int64_t replicated_total = 0; - for (const auto& [unit, entry] : declared) { - declared_total += entry->declared_bytes; - for (const PoolByteSpan& span : entry->spans) { - if (span.dst_unit_ordinal >= static_cast(num_dst)) { - return absl::InvalidArgumentError( - absl::StrCat("Byte span dst_unit_ordinal ", span.dst_unit_ordinal, - " exceeds the transfer's ", num_dst, - " destination units for tag ", PyStrRepr(plan_tag), - " (declared by ", PythonRepr(unit), ")")); - } - if (span.dst_block_index >= static_cast(dst_ids_g.size())) { - return absl::InvalidArgumentError( - absl::StrCat("Byte span destination index ", span.dst_block_index, - " exceeds the transfer's ", dst_ids_g.size(), - " destination blocks for tag ", PyStrRepr(plan_tag), - " (declared by ", PythonRepr(unit), ")")); - } - const int64_t src_end = span.src_offset_bytes + - (span.count - 1) * span.src_stride_bytes + - span.size_bytes; - if (src_end > src_live) { - return absl::InvalidArgumentError(absl::StrCat( - "Byte span exceeds its source block live bytes: end=", src_end, - ", live=", src_live, " (declared by ", PythonRepr(unit), ")")); - } - const int64_t dst_end = span.dst_offset_bytes + - (span.count - 1) * span.dst_stride_bytes + - span.size_bytes; - if (dst_end > dst_live) { - return absl::InvalidArgumentError(absl::StrCat( - "Byte span exceeds its destination block live bytes: end=", - dst_end, ", live=", dst_live, " (declared by ", PythonRepr(unit), - ")")); - } - expanded_repeats += span.count; - if (expanded_repeats > kMaxLiveSegments) { - return absl::InvalidArgumentError( - "Byte span plan exceeds the repeat expansion bound"); - } - for (int32_t repeat = 0; repeat < span.count; ++repeat) { - const int64_t start = - span.dst_offset_bytes + repeat * span.dst_stride_bytes; - if (span.dst_unit_ordinal < 0) { - for (size_t d = 0; d < num_dst; ++d) { - coverage_by_dst[d][span.dst_block_index].emplace_back( - start, start + span.size_bytes); - } - replicated_total += span.size_bytes; - } else { - coverage_by_dst[span.dst_unit_ordinal][span.dst_block_index] - .emplace_back(start, start + span.size_bytes); - } - covered_total += span.size_bytes; - } - } - } std::vector extents; - for (size_t d = 0; d < num_dst; ++d) { - // Single-destination plans keep the historical error strings; the - // destination suffix appears only for multi-destination plans. - const std::string dst_suffix = - num_dst > 1 ? absl::StrCat(" at destination unit ", d) : ""; - std::vector dst_extents; - for (size_t index = 0; index < dst_ids_g.size(); ++index) { - std::vector>& intervals = - coverage_by_dst[d][index]; - std::sort(intervals.begin(), intervals.end()); - if (intervals.empty()) { - return absl::InvalidArgumentError( - absl::StrCat("Destination block index ", index, - " has no declared coverage for tag ", - PyStrRepr(plan_tag), dst_suffix)); - } - int64_t covered_until = 0; - for (const auto& [start, end] : intervals) { - if (start != covered_until) { - const bool overlap = start < covered_until; - return absl::InvalidArgumentError(absl::StrCat( - "Declared byte spans have a destination coverage ", - overlap ? "overlap" : "gap", " at byte ", - std::min(start, covered_until), " of destination block index ", - index, dst_suffix)); - } - covered_until = end; - } - if (index != dst_ids_g.size() - 1 && covered_until != dst_live) { - return absl::InvalidArgumentError(absl::StrCat( - "Destination block index ", index, " is covered to ", - covered_until, " of ", dst_live, - " live bytes; only the final block may be partial", dst_suffix)); - } - dst_extents.push_back(covered_until); - } - if (d == 0) { - extents = std::move(dst_extents); - } else if (dst_extents != extents) { - // One extent vector per group travels to every receiver; sharded - // destinations must therefore cover the same byte prefix of every - // destination block (true for head-split caches by construction). - return absl::InvalidArgumentError(absl::StrCat( - "Destination units must share uniform coverage extents for tag ", - PyStrRepr(plan_tag), "; destination unit ", d, - " differs from destination unit 0")); - } - } - if (declared_total != covered_total + tag_clipped_bytes) { - if (tag_clipped_bytes == 0) { - return absl::InvalidArgumentError(absl::StrCat( - "Declared byte totals disagree with coverage: declared=", - declared_total, ", covered=", covered_total)); - } - return absl::InvalidArgumentError(absl::StrCat( - "Declared byte totals disagree with coverage: declared=", - declared_total, ", covered=", covered_total, - ", clipped=", tag_clipped_bytes)); - } - { - // Every destination covers extent_sum bytes; replicated spans are - // counted once in covered_total but reach all num_dst destinations. - int64_t extent_sum = 0; - for (int64_t extent : extents) extent_sum += extent; - if (static_cast(num_dst) * extent_sum != - covered_total + - static_cast(num_dst - 1) * replicated_total) { - return absl::InternalError("byte coverage accounting failed"); - } - } + absl::Status coverage_status = + ValidateSpanCoverage(declared, plan_tag, num_dst, dst_ids_g.size(), + src_live, dst_live, tag_clipped_bytes, &extents); + if (!coverage_status.ok()) return coverage_status; - // Bind ordinals and destination indices to physical ids and emit. - struct OrderedSpan { - const PoolByteSpan* span; - RaidenId unit; - const PoolSpanRegistration* entry; - }; - std::vector ordered_spans; - for (const auto& [unit, entry] : declared) { - for (const PoolByteSpan& span : entry->spans) { - ordered_spans.push_back(OrderedSpan{&span, unit, entry}); - } - } - std::stable_sort( - ordered_spans.begin(), ordered_spans.end(), - [](const OrderedSpan& a, const OrderedSpan& b) { - return std::tie(a.span->dst_block_index, a.span->dst_offset_bytes) < - std::tie(b.span->dst_block_index, b.span->dst_offset_bytes); - }); - std::map< - RaidenId, - std::map>, - RequestBlockRegistry::RaidenIdLess>, - RequestBlockRegistry::RaidenIdLess> - transfer_pairs_per_sender; - for (const OrderedSpan& ordered : ordered_spans) { - const PoolByteSpan& span = *ordered.span; - const RaidenId& src_unit = ordered.unit; - const int64_t src_block_id = - ordered.entry->block_ids[span.src_block_ordinal]; - const int64_t dst_block_id = dst_ids_g[span.dst_block_index]; - // A span without an ordinal replicates its chunks to every - // destination (entries differ only in dst_peer); a span with an - // ordinal reaches that destination only. - std::vector targets; - if (span.dst_unit_ordinal < 0) { - targets = request.dst_units; - } else { - targets.push_back(request.dst_units[span.dst_unit_ordinal]); - } - for (int32_t repeat = 0; repeat < span.count; ++repeat) { - const int64_t src_offset = - span.src_offset_bytes + repeat * span.src_stride_bytes; - const int64_t dst_offset = - span.dst_offset_bytes + repeat * span.dst_stride_bytes; - auto translated = - TranslateLiveCopy(precheck.src_segments, precheck.dst_segments, - src_offset, dst_offset, span.size_bytes); - if (!translated.ok()) return translated.status(); - emitted_chunks += static_cast(translated->size()) * - static_cast(targets.size()); - if (emitted_chunks > kMaxLiveSegments) { - return absl::InvalidArgumentError( - "Byte-span plan exceeds the live-region expansion bound"); - } - for (const LiveCopyChunk& chunk : *translated) { - for (const RaidenId& dst_unit_id : targets) { - ScheduleEntry schedule_entry; - schedule_entry.dst_peer = dst_peers.at(dst_unit_id); - schedule_entry.dst_shard_idx = 0; - schedule_entry.dst_offset_bytes = chunk.dst_physical; - schedule_entry.src_offset_bytes = chunk.src_physical; - schedule_entry.size_bytes = chunk.size; - schedule_entry.src_block_id = src_block_id; - schedule_entry.dst_block_id = dst_block_id; - schedule_entry.src_stride_bytes = 0; - schedule_entry.dst_stride_bytes = 0; - schedule_entry.count = 1; - schedule_entry.layer_idx = 0; - schedule_entry.pool_group = static_cast(group_idx); - schedules[src_unit].push_back(std::move(schedule_entry)); - } - } - } - auto& sender_pairs = transfer_pairs_per_sender[src_unit]; - for (const RaidenId& dst_unit_id : targets) { - sender_pairs[dst_unit_id].insert(std::make_tuple( - dst_peers.at(dst_unit_id), src_block_id, dst_block_id)); - } - } + TransferPairs transfer_pairs_per_sender; + absl::Status emit_status = + EmitSpanSchedule(declared, request.dst_units, dst_peers, dst_ids_g, + precheck, static_cast(group_idx), &schedules, + &emitted_chunks, &transfer_pairs_per_sender); + if (!emit_status.ok()) return emit_status; // Computed expected pushes for the receiver. std::map diff --git a/tpu_sync/kv_cache/reshard/reshard_service_test.cc b/tpu_sync/kv_cache/reshard/reshard_service_test.cc index 65b1310c..dd13c065 100644 --- a/tpu_sync/kv_cache/reshard/reshard_service_test.cc +++ b/tpu_sync/kv_cache/reshard/reshard_service_test.cc @@ -311,6 +311,67 @@ class ReshardStackTest : public ::testing::Test { return HandleController(req.SerializeAsString()); } + // One rank's request-global (v1) declaration of strided spans over one + // source block: a head slice of `size` bytes repeated `count` times per + // page, routed per destination unit. + struct StridedSpan { + int64_t src_offset; + int64_t dst_global_offset; + int64_t size; + int64_t src_stride; + int64_t dst_stride; + int32_t count; + int32_t dst_unit_ordinal; + }; + tpu_sync::rpc::ControllerResponse RegisterStridedGlobalSpans( + int rank, const std::string& req_id, int64_t uuid, int64_t src_block, + const std::vector& spans) { + tpu_sync::rpc::ControllerRequest req; + req.set_command( + tpu_sync::rpc::ControllerRequest::COMMAND_REGISTER_REQUEST_BLOCKS); + auto* block_req = req.mutable_register_request_blocks_request(); + block_req->set_req_id(req_id); + block_req->set_uuid(uuid); + *block_req->mutable_unit() = RaidenIdToProto(Unit(rank)); + block_req->add_block_ids(src_block); + auto* entry = block_req->add_pool_spans(); + entry->set_tag("fa"); + entry->add_block_ids(src_block); + int64_t declared = 0; + for (const StridedSpan& span : spans) { + auto* out = entry->add_spans(); + out->set_src_block_ordinal(0); + out->set_src_offset_bytes(span.src_offset); + out->set_dst_block_index(0); + out->set_dst_offset_bytes(span.dst_global_offset); + out->set_size_bytes(span.size); + out->set_src_stride_bytes(span.src_stride); + out->set_dst_stride_bytes(span.dst_stride); + out->set_count(span.count); + if (span.dst_unit_ordinal != -1) { + out->set_dst_unit_ordinal(span.dst_unit_ordinal); + } + declared += span.size * span.count; + } + entry->set_declared_bytes(declared); + entry->set_dst_space_version(1); + return HandleController(req.SerializeAsString()); + } + + // Head-split declarations for RegisterShardedPair(): a source page holds + // four tokens of two 128-byte heads; destination unit d takes head d of + // every token as one strided span per page. + void RegisterHeadSplitPages(const std::string& req_id, int64_t uuid) { + ASSERT_TRUE(RegisterStridedGlobalSpans(0, req_id, uuid, 3, + {{0, 0, 128, 256, 128, 4, 0}, + {128, 0, 128, 256, 128, 4, 1}}) + .success()); + ASSERT_TRUE(RegisterStridedGlobalSpans(1, req_id, uuid, 5, + {{0, 512, 128, 256, 128, 4, 0}, + {128, 512, 128, 256, 128, 4, 1}}) + .success()); + } + // Two sources (live 1024) feeding two head-shard destinations (live 512): // the TP2 shape where each destination receives a different half of // every source block. @@ -1472,6 +1533,108 @@ TEST_F(ReshardStackTest, TwoDestinationsShardedSpansRouteBytesPerReceiver) { } } +TEST_F(ReshardStackTest, StridedGlobalSpansStayOneEntryPerDestination) { + RegisterShardedPair(); + RegisterHeadSplitPages("req-st", 71); + tpu_sync::rpc::ControllerResponse resp = + Coordinate("req-st", 71, 2, {7, 9}, /*dst_skip=*/{}, /*num_dst=*/2); + ASSERT_TRUE(resp.success()) << resp.message(); + + // 2 concurrent arms, then 2 sender dispatches. + ASSERT_EQ(transport_.calls_.size(), 4u); + for (int i = 0; i < 2; ++i) { + tpu_sync::rpc::ControlRequest arm; + ASSERT_TRUE(arm.ParseFromString(transport_.calls_[i].second)); + const auto& start_req = arm.start_transfer_request(); + EXPECT_FALSE(start_req.is_sender()); + const bool is_dst0 = transport_.calls_[i].first == "10.0.0.2:9600"; + const std::string own_peer = is_dst0 ? "10.0.0.2:9400" : "10.0.0.2:9401"; + const int64_t own_src_offset = is_dst0 ? 0 : 128; + // Each sender contributes one strided entry: four tokens of this + // destination's head, not four entries. + ASSERT_EQ(start_req.shard_push_schedules_size(), 2); + for (const auto& keyed_schedule : start_req.shard_push_schedules()) { + ASSERT_EQ(keyed_schedule.second.entries_size(), 1); + const auto& entry = keyed_schedule.second.entries(0); + EXPECT_EQ(entry.dst_peer(), own_peer); + EXPECT_EQ(entry.src_offset_bytes(), own_src_offset); + EXPECT_EQ(entry.dst_offset_bytes(), 0); + EXPECT_EQ(entry.size_bytes(), 128); + EXPECT_EQ(entry.src_stride_bytes(), 256); + EXPECT_EQ(entry.dst_stride_bytes(), 128); + EXPECT_EQ(entry.count(), 4); + } + ASSERT_EQ(start_req.pool_groups_size(), 1); + EXPECT_EQ(start_req.pool_groups(0).expected_pushes(), 2); + ASSERT_EQ(start_req.pool_groups(0).dst_expected_extent_bytes_size(), 2); + EXPECT_EQ(start_req.pool_groups(0).dst_expected_extent_bytes(0), 512); + EXPECT_EQ(start_req.pool_groups(0).dst_expected_extent_bytes(1), 512); + } + for (int i = 2; i < 4; ++i) { + tpu_sync::rpc::ControlRequest dispatch; + ASSERT_TRUE(dispatch.ParseFromString(transport_.calls_[i].second)); + const auto& start_req = dispatch.start_transfer_request(); + EXPECT_TRUE(start_req.is_sender()); + ASSERT_EQ(start_req.shard_push_schedules_size(), 1); + const auto& schedule = start_req.shard_push_schedules().at(0); + ASSERT_EQ(schedule.entries_size(), 2); + for (const auto& entry : schedule.entries()) { + const int64_t expected_src = + entry.dst_peer() == "10.0.0.2:9400" ? 0 : 128; + EXPECT_EQ(entry.src_offset_bytes(), expected_src); + EXPECT_EQ(entry.src_stride_bytes(), 256); + EXPECT_EQ(entry.dst_stride_bytes(), 128); + EXPECT_EQ(entry.count(), 4); + } + EXPECT_NE(schedule.entries(0).dst_peer(), schedule.entries(1).dst_peer()); + } +} + +TEST_F(ReshardStackTest, StridedGlobalSpansClipByWholePages) { + RegisterShardedPair(); + RegisterHeadSplitPages("req-sc", 72); + // Skipping the first destination page drops rank 0's page whole and + // re-bases rank 1's page to destination index 0. + tpu_sync::rpc::ControllerResponse resp = + Coordinate("req-sc", 72, 2, {9}, /*dst_skip=*/{512}, /*num_dst=*/2); + ASSERT_TRUE(resp.success()) << resp.message(); + // 2 arms, then the one sender with surviving spans. + ASSERT_EQ(transport_.calls_.size(), 3u); + tpu_sync::rpc::ControlRequest dispatch; + ASSERT_TRUE(dispatch.ParseFromString(transport_.calls_[2].second)); + const auto& start_req = dispatch.start_transfer_request(); + EXPECT_TRUE(start_req.is_sender()); + EXPECT_EQ(transport_.calls_[2].first, "10.0.0.1:9102"); + const auto& schedule = start_req.shard_push_schedules().at(0); + ASSERT_EQ(schedule.entries_size(), 2); + for (const auto& entry : schedule.entries()) { + EXPECT_EQ(entry.src_block_id(), 5); + EXPECT_EQ(entry.dst_block_id(), 9); + EXPECT_EQ(entry.dst_offset_bytes(), 0); + EXPECT_EQ(entry.count(), 4); + } + for (int i = 0; i < 2; ++i) { + tpu_sync::rpc::ControlRequest arm; + ASSERT_TRUE(arm.ParseFromString(transport_.calls_[i].second)); + const auto& group = arm.start_transfer_request().pool_groups(0); + EXPECT_EQ(group.expected_pushes(), 1); + ASSERT_EQ(group.dst_expected_extent_bytes_size(), 1); + EXPECT_EQ(group.dst_expected_extent_bytes(0), 512); + } +} + +TEST_F(ReshardStackTest, StridedGlobalSpansMustStayInsideOnePage) { + RegisterShardedPair(); + // Eight repeats of 128 bytes at stride 128 span 1024 bytes: two + // destination pages. + tpu_sync::rpc::ControllerResponse resp = RegisterStridedGlobalSpans( + 0, "req-sx", 73, 3, {{0, 0, 128, 128, 128, 8, -1}}); + ASSERT_TRUE(resp.success()) << resp.message(); + resp = Coordinate("req-sx", 73, 1, {7, 9}, /*dst_skip=*/{}, /*num_dst=*/2); + EXPECT_FALSE(resp.success()); + EXPECT_THAT(resp.message(), HasSubstr("inside one destination page")); +} + TEST_F(ReshardStackTest, ShardedSpanOrdinalOutOfRangeFailsClosed) { RegisterShardedPair(); // Ordinal 2 with two destination units: accepted by the registry (the