40 lines
1.6 KiB
Python
40 lines
1.6 KiB
Python
|
|
from __future__ import annotations
|
||
|
|
from dataclasses import dataclass
|
||
|
|
import numpy as np
|
||
|
|
|
||
|
|
|
||
|
|
@dataclass(frozen=True)
|
||
|
|
class ChannelView:
|
||
|
|
"""Доступ к компонентам Y, Cb, Cr в массиве YCbCr."""
|
||
|
|
|
||
|
|
_MAP = {"y": 0, "cb": 1, "cr": 2}
|
||
|
|
|
||
|
|
@staticmethod
|
||
|
|
def _idx(name: str) -> int:
|
||
|
|
"""Вернуть индекс канала."""
|
||
|
|
k = name.strip().lower()
|
||
|
|
if k not in ChannelView._MAP:
|
||
|
|
raise ValueError("Канал должен быть Y, Cb или Cr.")
|
||
|
|
return ChannelView._MAP[k]
|
||
|
|
|
||
|
|
@staticmethod
|
||
|
|
def get(ycbcr: np.ndarray, channel: str) -> np.ndarray:
|
||
|
|
"""Вернуть 2D-компоненту выбранного канала."""
|
||
|
|
if ycbcr.ndim != 3 or ycbcr.shape[-1] != 3:
|
||
|
|
raise ValueError("Ожидается массив HxWx3.")
|
||
|
|
i = ChannelView._idx(channel)
|
||
|
|
return ycbcr[..., i]
|
||
|
|
|
||
|
|
@staticmethod
|
||
|
|
def set(ycbcr: np.ndarray, channel: str, component: np.ndarray) -> np.ndarray:
|
||
|
|
"""Вернуть копию YCbCr с заменённым каналом."""
|
||
|
|
if ycbcr.ndim != 3 or ycbcr.shape[-1] != 3:
|
||
|
|
raise ValueError("Ожидается массив HxWx3.")
|
||
|
|
if component.shape != ycbcr.shape[:2]:
|
||
|
|
raise ValueError("Размер компоненты не совпадает с HxW.")
|
||
|
|
out = ycbcr.copy()
|
||
|
|
i = ChannelView._idx(channel)
|
||
|
|
# Короткая отсечка и приведение типа
|
||
|
|
out[..., i] = np.clip(component, 0, 255).astype(np.uint8)
|
||
|
|
return out
|