MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / _conv_split

Function _conv_split

sat/vae_modules/cp_enc_dec.py:135–159  ·  view source on GitHub ↗
(input_, dim, kernel_size)

Source from the content-addressed store, hash-verified

133
134
135def _conv_split(input_, dim, kernel_size):
136 cp_world_size = get_context_parallel_world_size()
137
138 # Bypass the function if context parallel is 1
139 if cp_world_size == 1:
140 return input_
141
142 # print('in _conv_split, cp_rank:', cp_rank, 'input_size:', input_.shape)
143
144 cp_rank = get_context_parallel_rank()
145
146 dim_size = (input_.size()[dim] - kernel_size) // cp_world_size
147
148 if cp_rank == 0:
149 output = input_.transpose(dim, 0)[: dim_size + kernel_size].transpose(dim, 0)
150 else:
151 # output = input_.transpose(dim, 0)[cp_rank * dim_size + 1:(cp_rank + 1) * dim_size + kernel_size].transpose(dim, 0)
152 output = input_.transpose(dim, 0)[
153 cp_rank * dim_size + kernel_size : (cp_rank + 1) * dim_size + kernel_size
154 ].transpose(dim, 0)
155 output = output.contiguous()
156
157 # print('out _conv_split, cp_rank:', cp_rank, 'input_size:', output.shape)
158
159 return output
160
161
162def _conv_gather(input_, dim, kernel_size):

Callers 5

get_inputMethod · 0.90
encodeMethod · 0.90
decodeMethod · 0.90
forwardMethod · 0.70
backwardMethod · 0.70

Calls 2

Tested by

no test coverage detected