MCPcopy Create free account
hub / github.com/Agent-RL/ReCall / build_memory_reference_from_module

Function build_memory_reference_from_module

src/verl/utils/memory_buffer.py:100–113  ·  view source on GitHub ↗
(module: torch.nn.Module,
                                       memory_buffers: Dict[torch.dtype, MemoryBuffer],
                                       maintain_weight=True)

Source from the content-addressed store, hash-verified

98
99
100def build_memory_reference_from_module(module: torch.nn.Module,
101 memory_buffers: Dict[torch.dtype, MemoryBuffer],
102 maintain_weight=True):
103 start_index = {}
104 for dtype in memory_buffers.keys():
105 start_index[dtype] = 0
106 for name, param in sorted(module.named_parameters()):
107 memory_buffer = memory_buffers[param.dtype]
108 buffer = memory_buffer.get(shape=param.shape, start_index=start_index[param.dtype])
109 # need to increment start_index
110 start_index[param.dtype] += calc_padded_numel(param.shape, dtype)
111 if maintain_weight:
112 buffer.copy_(param.data)
113 param.data = buffer
114
115
116def build_memory_reference(weight_buffer_meta: Dict[str, Dict], memory_buffers: Dict[torch.dtype, MemoryBuffer]):

Callers 4

__init__Method · 0.85

Calls 3

calc_padded_numelFunction · 0.85
named_parametersMethod · 0.80
getMethod · 0.45

Tested by

no test coverage detected