CoolFace
Modelpublic

1ST-PLACE-WINNER/MiniMax-H3

sourceHugging Faceotherupdated 2d agoView on Hugging Face
0likes77downloads
parallel.py419 linesDownload Raw Back to video_vae
1# SPDX-License-Identifier: Apache-2.02# Parallel state and collective helpers for the MiniMax H3 visual VAE.3import os4import math5import torch6import torch.nn.functional as F7import torch.distributed as dist8from torch.autograd import Function9from torch.distributed import group, ReduceOp10 11 12def get_group_rank(group_size):13    global_rank = int(os.environ["RANK"])14    group_rank = global_rank % group_size15    return group_rank16 17 18_parallel_state = {}19 20# The torch.autograd.Function subclasses below keep their backward() methods21# to satisfy the autograd.Function contract; only the forward paths are22# exercised in this inference-only bundle.23 24 25def get_parallel_state():26    return _parallel_state27 28 29class _AllGather(Function):30    @staticmethod31    def forward(ctx, group, tensor):32        tensor = tensor.contiguous()33        ctx.group = group34        group_size = dist.get_world_size(group=group)35        out_tensor_list = [torch.empty_like(tensor) for _ in range(group_size)]36        dist.all_gather(out_tensor_list, tensor, group=group)37        return tuple(out_tensor_list)38 39    @staticmethod40    def backward(ctx, *grad_outputs):41        rank = dist.get_rank(group=ctx.group)42        gx = torch.empty_like(grad_outputs[rank])43        gx = gx.contiguous()44        grad_outputs = tuple(t.contiguous() for t in grad_outputs)45        dist.reduce_scatter(gx, list(grad_outputs), op=ReduceOp.SUM, group=ctx.group)46        return (None, gx)47 48 49@torch.compiler.disable50def all_gather(tensor, group=group.WORLD):51    return _AllGather.apply(group, tensor)52 53 54class _AllGatherVarShape(Function):55    @staticmethod56    def forward(ctx, group, tensor):57        tensor = tensor.contiguous()58        ctx.group = group59        ctx.original_shape = tensor.shape60 61        shape_info = torch.tensor(62            list(tensor.shape), dtype=torch.long, device=tensor.device63        )64 65        shape_list = [66            torch.empty_like(shape_info)67            for _ in range(dist.get_world_size(group=group))68        ]69        dist.all_gather(shape_list, shape_info, group=group)70 71        all_shapes = [tuple(shape_tensor.tolist()) for shape_tensor in shape_list]72        ctx.all_shapes = all_shapes73 74        flat_tensor = tensor.flatten()75        max_size = max(math.prod(s) for s in all_shapes)76 77        if flat_tensor.numel() < max_size:78            padded = torch.zeros(max_size, dtype=tensor.dtype, device=tensor.device)79            padded[: flat_tensor.numel()] = flat_tensor80            flat_tensor = padded81 82        gathered_flat = [torch.empty_like(flat_tensor) for _ in range(len(all_shapes))]83        dist.all_gather(gathered_flat, flat_tensor, group=group)84 85        return tuple(86            t[: math.prod(shape)].reshape(shape)87            for t, shape in zip(gathered_flat, all_shapes)88        )89 90    @staticmethod91    def backward(ctx, *grad_outputs):92        rank = dist.get_rank(group=ctx.group)93 94        grad_input = grad_outputs[rank]95        if grad_input is None:96            return None, torch.zeros(97                ctx.original_shape, device=next(iter(grad_outputs)).device98            )99 100        max_size = max(math.prod(shape) for shape in ctx.all_shapes)101        padded_grads = []102 103        for grad, shape in zip(grad_outputs, ctx.all_shapes):104            if grad is not None:105                flat_grad = grad.flatten()106            else:107                flat_grad = torch.zeros(108                    math.prod(shape),109                    dtype=grad_input.dtype,110                    device=grad_input.device,111                )112 113            if flat_grad.numel() < max_size:114                padded = torch.zeros(115                    max_size, dtype=flat_grad.dtype, device=flat_grad.device116                )117                padded[: flat_grad.numel()] = flat_grad118                padded_grads.append(padded)119            else:120                padded_grads.append(flat_grad)121 122        result_grad = torch.empty_like(padded_grads[0])123        dist.reduce_scatter(result_grad, padded_grads, op=ReduceOp.SUM, group=ctx.group)124 125        original_size = math.prod(ctx.original_shape)126        return None, result_grad[:original_size].reshape(ctx.original_shape)127 128 129@torch.compiler.disable130def all_gather_var_shape(tensor, group=group.WORLD):131    return _AllGatherVarShape.apply(group, tensor)132 133 134class _AllReduce(Function):135    @staticmethod136    def forward(ctx, _input, op, group):137        ctx.group = group138        ctx.op = op139        _input = _input.clone()140        dist.all_reduce(_input, op=op, group=group)141        return _input142 143    @staticmethod144    def backward(ctx, grad_output):145        grad_output = grad_output.clone()146        dist.all_reduce(grad_output, op=ctx.op, group=ctx.group)147        return grad_output, None, None148 149 150@torch.compiler.disable151def all_reduce(input_, op, group):152    return _AllReduce.apply(input_, op, group)153 154 155class _AlltoAllSingle(Function):156    @staticmethod157    def forward(ctx, group, input):158        ctx.group = group159 160        world_size = dist.get_world_size(group=group)161        if world_size == 1:162            return input163 164        input = input.contiguous()165        output = torch.empty_like(input)166        dist.all_to_all_single(167            output,168            input,169            group=group,170        )171        return output172 173    @staticmethod174    def backward(ctx, grad_output):175        return (None, _AlltoAllSingle.apply(ctx.group, grad_output))176 177 178@torch.compiler.disable179def all_to_all_single(input, group=group.WORLD):180    return _AlltoAllSingle.apply(group, input)181 182 183 184@torch.compiler.disable185def get_subseq(input, sp_size=None):186    if sp_size is None:187        state = get_parallel_state()188        if not state.get("sp_enabled", False):189            return input190        sp_size = state["sp_size"]191        sp_rank = state["sp_rank"]192    else:193        sp_rank = get_group_rank(sp_size)194 195    if sp_size == 1:196        return input197 198    if input.shape[1] % sp_size != 0:199        raise ValueError(200            f"Input shape {input.shape} is not divisible by sp_size {sp_size}"201        )202 203    return torch.chunk(input, sp_size, dim=1)[sp_rank]204 205 206@torch.compiler.disable207def gather_subseq(input, sp_size=None, local_process_group=None):208    if sp_size is None:209        state = get_parallel_state()210        if not state.get("sp_enabled", False):211            return input212        sp_size = state["sp_size"]213        local_process_group = state["sp_process_group"]214 215    if sp_size == 1:216        return input217 218    output = all_gather(input, group=local_process_group)219    output = torch.cat(output, dim=1)220    return output221 222 223@torch.compiler.disable224def all_to_all_4D(225    input: torch.tensor,226    scatter_idx: int = 2,227    gather_idx: int = 1,228    group=None,229):230    assert (231        input.dim() == 4232    ), f"input must be 4D tensor, got {input.dim()} and shape {input.shape}"233 234    if group is None:235        seq_world_size = 1236    else:237        seq_world_size = dist.get_world_size(group)238 239    if seq_world_size == 1:240        return input241 242    if scatter_idx == 2 and gather_idx == 1:243        bs, shard_seqlen, hc, hs = input.shape244        seqlen = shard_seqlen * seq_world_size245        shard_hc = hc // seq_world_size246 247        input_t = (248            input.reshape(bs, shard_seqlen, seq_world_size, shard_hc, hs)249            .transpose(0, 2)250            .contiguous()251        )252 253        output = all_to_all_single(input_t, group=group)254        output = output.reshape(seqlen, bs, shard_hc, hs)255        output = output.transpose(0, 1).contiguous().reshape(bs, seqlen, shard_hc, hs)256        return output257 258    elif scatter_idx == 1 and gather_idx == 2:259        bs, seqlen, shard_hc, hs = input.shape260        hc = shard_hc * seq_world_size261        shard_seqlen = seqlen // seq_world_size262 263        input_t = (264            input.reshape(bs, seq_world_size, shard_seqlen, shard_hc, hs)265            .transpose(0, 3)266            .transpose(0, 1)267            .contiguous()268            .reshape(seq_world_size, shard_hc, shard_seqlen, bs, hs)269        )270 271        output = all_to_all_single(input_t, group=group)272        output = output.reshape(hc, shard_seqlen, bs, hs)273        output = output.transpose(0, 2).contiguous().reshape(bs, shard_seqlen, hc, hs)274        return output275    else:276        raise RuntimeError("scatter_idx must be 1 or 2 and gather_idx must be 1 or 2")277 278 279 280@torch.compiler.disable281def exchange_borders(282    input_, padding, pad_mode, sp_rank, sp_size, group, dim=-1, async_op=False283):284    if async_op and input_.requires_grad:285        raise ValueError("async_op is not supported backward, check previous commits")286 287    slice_indices = [slice(None)] * input_.ndim288    slice_indices[dim] = slice(None, padding)289    first_tensor = input_[tuple(slice_indices)].contiguous()290 291    slice_indices[dim] = slice(-padding, None)292    last_tensor = input_[tuple(slice_indices)].contiguous()293 294    if async_op:295        first_borders = [torch.empty_like(first_tensor) for _ in range(sp_size)]296        last_borders = [torch.empty_like(last_tensor) for _ in range(sp_size)]297 298        handle_first = dist.all_gather(299            first_borders, first_tensor, group=group, async_op=True300        )301        handle_last = dist.all_gather(302            last_borders, last_tensor, group=group, async_op=True303        )304    else:305        first_borders = all_gather(first_tensor, group=group)306        last_borders = all_gather(last_tensor, group=group)307 308    if dim < 0:309        pad_dim = -1 - dim310    else:311        pad_dim = input_.ndim - 1 - dim312 313    pad_size = [0] * ((input_.ndim - 2) * 2)314    pad_size[pad_dim * 2] = padding315    pad_size[pad_dim * 2 + 1] = padding316    output = F.pad(input_, pad_size, mode=pad_mode)317 318    slice_indices = [slice(None)] * input_.ndim319    slice_indices[dim] = slice(-padding, None)320 321    if async_op:322        handle_first.wait()323 324    if sp_rank < sp_size - 1:325        output[tuple(slice_indices)] = first_borders[sp_rank + 1]326    else:327        output[tuple(slice_indices)] += first_borders[0] * 0.0328 329    slice_indices = [slice(None)] * input_.ndim330    slice_indices[dim] = slice(None, padding)331 332    if async_op:333        handle_last.wait()334 335    if sp_rank > 0:336        output[tuple(slice_indices)] = last_borders[sp_rank - 1]337    else:338        output[tuple(slice_indices)] += last_borders[sp_size - 1] * 0.0339 340    return output341 342 343@torch.compiler.disable344def exchange_strides(345    input_, pad_mode, sp_rank, sp_size, group, dim=-1, async_op=False346):347    if async_op and input_.requires_grad:348        raise ValueError("async_op is not supported backward, check previous commits")349 350    if dim not in [-1, -2]:351        raise ValueError("dim must be -1 (W) or -2 (H) for exchange_strides")352 353    if dim == -1:354        if input_.ndim == 5:355            input_ = F.pad(input_, (0, 0, 0, 1, 0, 0), mode=pad_mode)356        elif input_.ndim == 4:357            input_ = F.pad(input_, (0, 0, 0, 1), mode=pad_mode)358        else:359            raise ValueError(f"Input must have 4 or 5 dimensions, got {input_.ndim}")360 361        left_border = input_[..., :1].contiguous()362 363        if async_op:364            left_borders = [torch.empty_like(left_border) for _ in range(sp_size)]365            handle = dist.all_gather(366                left_borders, left_border, group=group, async_op=True367            )368        else:369            left_borders = all_gather(left_border, group=group)370 371        if input_.ndim == 5:372            output = F.pad(input_, (0, 1, 0, 0, 0, 0), mode=pad_mode)373        elif input_.ndim == 4:374            output = F.pad(input_, (0, 1, 0, 0), mode=pad_mode)375        else:376            raise ValueError(f"Input must have 4 or 5 dimensions, got {input_.ndim}")377 378        if async_op:379            handle.wait()380 381        if sp_rank != sp_size - 1:382            output[..., -1:] = left_borders[sp_rank + 1]383        else:384            output[..., -1:] += left_borders[0] * 0.0385    else:386        if input_.ndim == 5:387            input_ = F.pad(input_, (0, 1, 0, 0, 0, 0), mode=pad_mode)388        elif input_.ndim == 4:389            input_ = F.pad(input_, (0, 1, 0, 0), mode=pad_mode)390        else:391            raise ValueError(f"Input must have 4 or 5 dimensions, got {input_.ndim}")392 393        top_border = input_[..., :1, :].contiguous()394 395        if async_op:396            top_borders = [torch.empty_like(top_border) for _ in range(sp_size)]397            handle = dist.all_gather(398                top_borders, top_border, group=group, async_op=True399            )400        else:401            top_borders = all_gather(top_border, group=group)402 403        if input_.ndim == 5:404            output = F.pad(input_, (0, 0, 0, 1, 0, 0), mode=pad_mode)405        elif input_.ndim == 4:406            output = F.pad(input_, (0, 0, 0, 1), mode=pad_mode)407        else:408            raise ValueError(f"Input must have 4 or 5 dimensions, got {input_.ndim}")409 410        if async_op:411            handle.wait()412 413        if sp_rank != sp_size - 1:414            output[..., -1:, :] = top_borders[sp_rank + 1]415        else:416            output[..., -1:, :] += top_borders[0] * 0.0417 418    return output419