Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
31 commits
Select commit Hold shift + click to select a range
96eecb4
Add a deterministic_random reward type
fzyzcjy Jun 22, 2026
2c1d4a7
Add an inplace_modify_args context manager
fzyzcjy Jun 22, 2026
7b00c75
Add fault-tolerance support tweaks to shared utilities
fzyzcjy Jun 22, 2026
635aa50
Preserve the process-group backend across reload
fzyzcjy Jun 22, 2026
7a33a63
Add fault-tolerance foundation utilities
fzyzcjy Jun 22, 2026
8dc8ddd
Add structured logfmt logging helper
fzyzcjy Jun 22, 2026
96d220f
Add a Clock abstraction with a fake clock for tests
fzyzcjy Jun 22, 2026
a005a2f
Add a fault injector test utility
fzyzcjy Jun 22, 2026
8365971
Add control-server data models
fzyzcjy Jun 22, 2026
e308821
Add a cell health checker and heartbeat utilities
fzyzcjy Jun 22, 2026
4a52fcf
Add the fault-tolerance dependency, CI label, and logger-config setup
fzyzcjy Jun 22, 2026
2ed9fcf
Always reconnect rollout engines on weight-update setup
fzyzcjy Jun 22, 2026
a838a0b
Add a fault-injection RPC to train actors
fzyzcjy Jun 22, 2026
c7798b5
Delay splitting train data by DP until actor-side processing
fzyzcjy Jul 8, 2026
d0b41fc
Add a deterministic NCCL backend for order-stable collectives
fzyzcjy Jun 22, 2026
a48d03f
Add a per-process identity helper
fzyzcjy Jun 22, 2026
3e85a86
Add structured event models keyed by per-process identity
fzyzcjy Jun 22, 2026
dc11be3
Add structured event logging keyed by per-process identity
fzyzcjy Jun 22, 2026
691e4a7
Add event-log snapshot and restore checkpointing
fzyzcjy Jun 22, 2026
5a91150
Log training metrics as MetricEvents through the event logger
fzyzcjy Jun 22, 2026
34fb714
Add the witness id allocator
fzyzcjy Jun 22, 2026
fd3bc9c
Trace witness ids through the model via injected witness parameters
fzyzcjy Jun 22, 2026
85f120b
Add event-log checksum-consistency analysis rules
fzyzcjy Jun 22, 2026
5b5c797
Add an event-log witness-tracing analysis rule
fzyzcjy Jun 22, 2026
493ef68
Add the event-log analyzer that applies analysis rules
fzyzcjy Jun 22, 2026
81b95c4
Add dump and inference-engine-checksum comparison helpers for FT tests
fzyzcjy Jun 22, 2026
71d38e2
Add metric comparison helpers for FT tests
fzyzcjy Jun 22, 2026
c043ccf
Add reconfiguration assertions for fault-tolerance tests
fzyzcjy Jun 22, 2026
932660e
Relocate GroupInfo into shared process-group utilities
fzyzcjy Jun 22, 2026
69c50b9
Add _TensorViewCodec for storage-deduplicated tensor serialization
fzyzcjy Jun 22, 2026
956b4b1
Merge remote-tracking branch 'origin/main' into tom/pr_chain/trainer_…
fzyzcjy Jul 10, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 52 additions & 0 deletions miles/backends/megatron_utils/checkpoint_transfer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
import torch


class _TensorViewCodec:
"""Encode tensors as (unique_storages, view_metas) and decode back.

Many input tensors may share underlying storage (e.g. Megatron
distributed-optimizer grad buckets). `encode` dedups by storage data_ptr —
each unique storage is wrapped once as a uint8 tensor (no copy), plus a
per-input view_meta record (storage_id, dtype, shape, stride,
storage_offset). `decode` reconstructs the original views with
`as_strided` over the storage bytes reinterpreted at the original dtype.
"""

@staticmethod
def encode(tensors: list[torch.Tensor]) -> tuple[list[torch.Tensor], list[dict]]:
storage_id_by_key: dict[tuple[torch.device, int], int] = {}
unique_storages: list[torch.Tensor] = []
view_metas: list[dict] = []
for t in tensors:
storage = t.untyped_storage()
key = (t.device, storage.data_ptr())
if key not in storage_id_by_key:
storage_id_by_key[key] = len(unique_storages)
# Wrap full storage as uint8 tensor (no copy, shares memory).
unique_storages.append(torch.tensor(storage, dtype=torch.uint8, device=t.device))
view_metas.append(
{
"storage_id": storage_id_by_key[key],
"dtype": t.dtype,
"shape": tuple(t.shape),
"stride": tuple(t.stride()),
"storage_offset": t.storage_offset(),
}
)
return unique_storages, view_metas

@staticmethod
def decode(unique_storages: list[torch.Tensor], view_metas: list[dict]) -> list[torch.Tensor]:
tensors: list[torch.Tensor] = []
for vm in view_metas:
storage_t = unique_storages[vm["storage_id"]] # uint8 view of received storage
# Reinterpret bytes as the original dtype, then apply stride/offset.
dtype_view = storage_t.view(vm["dtype"])
view = torch.as_strided(
dtype_view,
size=vm["shape"],
stride=vm["stride"],
storage_offset=vm["storage_offset"],
)
tensors.append(view)
Comment on lines +41 to +51

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

medium

If the untyped storage size is not a multiple of the element size of vm["dtype"] (which can happen with mixed-dtype storages or padded storages), calling storage_t.view(vm["dtype"]) will raise a RuntimeError. To prevent this, slice storage_t to a multiple of the element size before reinterpreting its dtype.

Suggested change
for vm in view_metas:
storage_t = unique_storages[vm["storage_id"]] # uint8 view of received storage
# Reinterpret bytes as the original dtype, then apply stride/offset.
dtype_view = storage_t.view(vm["dtype"])
view = torch.as_strided(
dtype_view,
size=vm["shape"],
stride=vm["stride"],
storage_offset=vm["storage_offset"],
)
tensors.append(view)
for vm in view_metas:
storage_t = unique_storages[vm["storage_id"]] # uint8 view of received storage
# Reinterpret bytes as the original dtype, then apply stride/offset.
element_size = torch.tensor([], dtype=vm["dtype"]).element_size()
num_bytes = (storage_t.numel() // element_size) * element_size
dtype_view = storage_t[:num_bytes].view(vm["dtype"])
view = torch.as_strided(
dtype_view,
size=vm["shape"],
stride=vm["stride"],
storage_offset=vm["storage_offset"],
)
tensors.append(view)

return tensors
Loading
Loading