File size: 4,469 Bytes
7faaef2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
"""studio.rigging.rigid_piece — rigid-piece backward-warp engine (Session 3.5).

Each piece rotates rigidly around its bone joint's position by that joint's
world rotation (computed via the same FK pass the LBS engine uses). Pixels are
sampled by backward warp: for each output pixel, the inverse transform maps to
a source pixel; if the source pixel belongs to this piece in the rest parts.png,
its color is copied.

Why backward-warp instead of forward-splat:
  - No gaps when a piece under-samples its output region after rotation.
  - z-order is the natural occlusion model — paint back-to-front, higher z wins.
  - Nearest-neighbor membership check matches sprite-art pixel discipline.
"""
from __future__ import annotations

from typing import Iterable, Mapping, Tuple

import numpy as np

from pixel_cursor.rigging import JOINT_INDEX, Skeleton

from .deform import forward_kinematics
from .parts import PartsMask


def deform_sprite_rigid(
    image: np.ndarray,
    rest: Skeleton,
    parts: PartsMask,
    *,
    local_rotations: Mapping[str, float] | None = None,
    include_kinds: Tuple[str, ...] = ("body",),
    exclude_pieces: Iterable[str] = (),
) -> Tuple[np.ndarray, np.ndarray]:
    """Render a posed configuration via rigid-piece backward warp.

    Args:
        image: source sprite, shape (H, W, 3 or 4) uint8.
        rest:  rest-pose skeleton (image_shape must equal image's H, W).
        parts: piece-mask sidecar (id_map shape must equal image's H, W).
        local_rotations: optional dict mapping joint_name -> local rotation (rad).
        include_kinds: piece.kind values to render. Default ('body',) — attachments
            are handled by the composite layer (Session 4).
        exclude_pieces: piece names to skip (e.g., ('sword',) for swordless render).

    Returns:
        (rgba (H, W, 4) uint8, mask (H, W) bool)

    The output is composed back-to-front by z; higher z overwrites lower. Every
    output pixel is a pixel-perfect copy of exactly one source pixel (or fully
    transparent) — this is the 'rigid' invariant.
    """
    H, W = parts.image_shape
    if image.shape[:2] != (H, W):
        raise ValueError(
            f"image shape {image.shape[:2]} != parts.image_shape {(H, W)}"
        )
    if rest.image_shape != (H, W):
        raise ValueError(
            f"rest.image_shape {rest.image_shape} != parts.image_shape {(H, W)}"
        )

    if image.ndim == 3 and image.shape[-1] == 4:
        rgba = image
    elif image.ndim == 3 and image.shape[-1] == 3:
        rgba = np.concatenate(
            [image, np.full((H, W, 1), 255, dtype=np.uint8)], axis=-1
        )
    elif image.ndim == 2:
        gray = image
        rgba = np.stack(
            [gray, gray, gray, np.full_like(gray, 255)], axis=-1
        )
    else:
        raise ValueError(f"unsupported image shape {image.shape}")

    deformed_pos, world_rot = forward_kinematics(rest, local_rotations or {})

    out_rgba = np.zeros((H, W, 4), dtype=np.uint8)
    out_mask = np.zeros((H, W), dtype=bool)

    yy, xx = np.mgrid[0:H, 0:W].astype(np.float32)

    exclude = set(exclude_pieces)
    selected = [
        p for p in parts.pieces
        if p.kind in include_kinds and p.name not in exclude
    ]
    selected.sort(key=lambda p: p.z)

    for piece in selected:
        bone_idx = JOINT_INDEX[piece.bone]
        pivot_rest_y = float(rest.positions[bone_idx, 0])
        pivot_rest_x = float(rest.positions[bone_idx, 1])
        pivot_posed_y = float(deformed_pos[bone_idx, 0])
        pivot_posed_x = float(deformed_pos[bone_idx, 1])
        theta = float(world_rot[bone_idx])

        c = np.cos(-theta)
        s = np.sin(-theta)

        dy = yy - pivot_posed_y
        dx = xx - pivot_posed_x
        src_y = pivot_rest_y + c * dy - s * dx
        src_x = pivot_rest_x + s * dy + c * dx

        sy_i = np.round(src_y).astype(np.int32)
        sx_i = np.round(src_x).astype(np.int32)
        in_bounds = (sy_i >= 0) & (sy_i < H) & (sx_i >= 0) & (sx_i < W)
        sy_safe = np.clip(sy_i, 0, H - 1)
        sx_safe = np.clip(sx_i, 0, W - 1)

        membership = parts.id_map[sy_safe, sx_safe] == piece.piece_id
        src_alpha = rgba[sy_safe, sx_safe, 3] > 0
        valid = in_bounds & membership & src_alpha
        if not valid.any():
            continue
        out_rgba[valid] = rgba[sy_safe[valid], sx_safe[valid]]
        out_mask |= valid

    return out_rgba, out_mask


__all__ = ["deform_sprite_rigid"]