CoolFace
Apppublic

siran002/meme-generator

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
utils.py449 linesDownload Raw Back to root
1import asyncio2import hashlib3import inspect4import math5import random6import time7from dataclasses import dataclass8from enum import Enum9from functools import partial, wraps10from io import BytesIO11from typing import (12    TYPE_CHECKING,13    Any,14    Callable,15    Coroutine,16    List,17    Literal,18    Optional,19    Protocol,20    Tuple,21    TypeVar,22)23 24import httpx25from PIL.Image import Image as IMG26from pil_utils import BuildImage, Text2Image27from pil_utils.types import ColorType, FontStyle, FontWeight28from typing_extensions import ParamSpec29 30from .config import meme_config31from .exception import MemeGeneratorException32 33if TYPE_CHECKING:34    from .meme import Meme35 36P = ParamSpec("P")37R = TypeVar("R")38 39 40def run_sync(call: Callable[P, R]) -> Callable[P, Coroutine[None, None, R]]:41    """一个用于包装 sync function 为 async function 的装饰器42    参数:43        call: 被装饰的同步函数44    """45 46    @wraps(call)47    async def _wrapper(*args: P.args, **kwargs: P.kwargs) -> R:48        loop = asyncio.get_running_loop()49        pfunc = partial(call, *args, **kwargs)50        result = await loop.run_in_executor(None, pfunc)51        return result52 53    return _wrapper54 55 56def is_coroutine_callable(call: Callable[..., Any]) -> bool:57    """检查 call 是否是一个 callable 协程函数"""58    if inspect.isroutine(call):59        return inspect.iscoroutinefunction(call)60    if inspect.isclass(call):61        return False62    func_ = getattr(call, "__call__", None)63    return inspect.iscoroutinefunction(func_)64 65 66def save_gif(frames: List[IMG], duration: float) -> BytesIO:67    output = BytesIO()68    frames[0].save(69        output,70        format="GIF",71        save_all=True,72        append_images=frames[1:],73        duration=duration * 1000,74        loop=0,75        disposal=2,76        optimize=False,77    )78 79    # 没有超出最大大小,直接返回80    nbytes = output.getbuffer().nbytes81    if nbytes <= meme_config.gif.gif_max_size * 10**6:82        return output83 84    # 超出最大大小,帧数超出最大帧数时,缩减帧数85    n_frames = len(frames)86    gif_max_frames = meme_config.gif.gif_max_frames87    if n_frames > gif_max_frames:88        index = range(n_frames)89        ratio = n_frames / gif_max_frames90        index = (int(i * ratio) for i in range(gif_max_frames))91        new_duration = duration * ratio92        new_frames = [frames[i] for i in index]93        return save_gif(new_frames, new_duration)94 95    # 超出最大大小,帧数没有超出最大帧数时,缩小尺寸96    new_frames = [97        frame.resize((int(frame.width * 0.9), int(frame.height * 0.9)))98        for frame in frames99    ]100    return save_gif(new_frames, duration)101 102 103class Maker(Protocol):104    def __call__(self, img: BuildImage) -> BuildImage:105        ...106 107 108class GifMaker(Protocol):109    def __call__(self, i: int) -> Maker:110        ...111 112 113def get_avg_duration(image: IMG) -> float:114    if not getattr(image, "is_animated", False):115        return 0116    total_duration = 0117    for i in range(image.n_frames):118        image.seek(i)119        total_duration += image.info["duration"]120    return total_duration / image.n_frames121 122 123def split_gif(image: IMG) -> List[IMG]:124    frames: List[IMG] = []125 126    update_mode = "full"127    for i in range(image.n_frames):128        image.seek(i)129        if image.tile:  # type: ignore130            update_region = image.tile[0][1][2:]  # type: ignore131            if update_region != image.size:132                update_mode = "partial"133                break134 135    last_frame: Optional[IMG] = None136    for i in range(image.n_frames):137        image.seek(i)138        frame = image.copy()139        if update_mode == "partial" and last_frame:140            frame = last_frame.copy().paste(frame)141        frames.append(frame)142    image.seek(0)143    if image.info.__contains__("transparency"):144        frames[0].info["transparency"] = image.info["transparency"]145    return frames146 147 148def make_jpg_or_gif(149    img: BuildImage, func: Maker, keep_transparency: bool = False150) -> BytesIO:151    """152    制作静图或者动图153    :params154      * ``img``: 输入图片155      * ``func``: 图片处理函数,输入img,返回处理后的图片156      * ``keep_transparency``: 传入gif时,是否保留该gif的透明度157    """158    image = img.image159    if not getattr(image, "is_animated", False):160        return func(img).save_jpg()161    else:162        frames = split_gif(image)163        duration = get_avg_duration(image) / 1000164        frames = [func(BuildImage(frame)).image for frame in frames]165        if keep_transparency:166            image.seek(0)167            if image.info.__contains__("transparency"):168                frames[0].info["transparency"] = image.info["transparency"]169        return save_gif(frames, duration)170 171 172def make_png_or_gif(173    img: BuildImage, func: Maker, keep_transparency: bool = False174) -> BytesIO:175    """176    制作静图或者动图177    :params178      * ``img``: 输入图片179      * ``func``: 图片处理函数,输入img,返回处理后的图片180      * ``keep_transparency``: 传入gif时,是否保留该gif的透明度181    """182    image = img.image183    if not getattr(image, "is_animated", False):184        return func(img).save_png()185    else:186        frames = split_gif(image)187        duration = get_avg_duration(image) / 1000188        frames = [func(BuildImage(frame)).image for frame in frames]189        if keep_transparency:190            image.seek(0)191            if image.info.__contains__("transparency"):192                frames[0].info["transparency"] = image.info["transparency"]193        return save_gif(frames, duration)194 195 196class FrameAlignPolicy(Enum):197    """198    要叠加的gif长度大于基准gif时,是否延长基准gif长度以对齐两个gif199    """200 201    no_extend = 0202    """不延长"""203    extend_first = 1204    """延长第一帧"""205    extend_last = 2206    """延长最后一帧"""207    extend_loop = 3208    """以循环方式延长"""209 210 211def make_gif_or_combined_gif(212    img: BuildImage,213    maker: GifMaker,214    frame_num: int,215    duration: float,216    frame_align: FrameAlignPolicy = FrameAlignPolicy.no_extend,217    input_based: bool = False,218    keep_transparency: bool = False,219) -> BytesIO:220    """221    使用静图或动图制作gif222    :params223      * ``img``: 输入图片,如头像224      * ``maker``: 图片处理函数生成,传入第几帧,返回对应的图片处理函数225      * ``frame_num``: 目标gif的帧数226      * ``duration``: 相邻帧之间的时间间隔,单位为秒227      * ``frame_align``: 要叠加的gif长度大于基准gif时,gif长度对齐方式228      * ``input_based``: 是否以输入gif为基准合成gif,默认为`False`,即以目标gif为基准229      * ``keep_transparency``: 传入gif时,是否保留该gif的透明度230    """231    image = img.image232    if not getattr(image, "is_animated", False):233        return save_gif([maker(i)(img).image for i in range(frame_num)], duration)234 235    frame_num_in = image.n_frames236    duration_in = get_avg_duration(image) / 1000237    total_duration_in = frame_num_in * duration_in238    total_duration = frame_num * duration239 240    if input_based:241        frame_num_base = frame_num_in242        frame_num_fit = frame_num243        duration_base = duration_in244        duration_fit = duration245        total_duration_base = total_duration_in246        total_duration_fit = total_duration247    else:248        frame_num_base = frame_num249        frame_num_fit = frame_num_in250        duration_base = duration251        duration_fit = duration_in252        total_duration_base = total_duration253        total_duration_fit = total_duration_in254 255    frame_idxs: List[int] = list(range(frame_num_base))256    diff_duration = total_duration_fit - total_duration_base257    diff_num = int(diff_duration / duration_base)258 259    if diff_duration >= duration_base:260        if frame_align == FrameAlignPolicy.extend_first:261            frame_idxs = [0] * diff_num + frame_idxs262 263        elif frame_align == FrameAlignPolicy.extend_last:264            frame_idxs += [frame_num_base - 1] * diff_num265 266        elif frame_align == FrameAlignPolicy.extend_loop:267            frame_num_total = frame_num_base268            # 重复基准gif,直到两个gif总时长之差在1个间隔以内,或总帧数超出最大帧数269            while frame_num_total + frame_num_base <= meme_config.gif.gif_max_frames:270                frame_num_total += frame_num_base271                frame_idxs += list(range(frame_num_base))272                multiple = round(frame_num_total * duration_base / total_duration_fit)273                if (274                    math.fabs(275                        total_duration_fit * multiple - frame_num_total * duration_base276                    )277                    <= duration_base278                ):279                    break280 281    frames: List[IMG] = []282    frame_idx_fit = 0283    time_start = 0284    for i, idx in enumerate(frame_idxs):285        while frame_idx_fit < frame_num_fit:286            if (287                frame_idx_fit * duration_fit288                <= i * duration_base - time_start289                < (frame_idx_fit + 1) * duration_fit290            ):291                if input_based:292                    idx_in = idx293                    idx_maker = frame_idx_fit294                else:295                    idx_in = frame_idx_fit296                    idx_maker = idx297 298                func = maker(idx_maker)299                image.seek(idx_in)300                frames.append(func(BuildImage(image.copy())).image)301                break302            else:303                frame_idx_fit += 1304                if frame_idx_fit >= frame_num_fit:305                    frame_idx_fit = 0306                    time_start += total_duration_fit307 308    if keep_transparency:309        image.seek(0)310        if image.info.__contains__("transparency"):311            frames[0].info["transparency"] = image.info["transparency"]312 313    return save_gif(frames, duration)314 315 316async def translate(text: str, lang_from: str = "auto", lang_to: str = "zh") -> str:317    appid = meme_config.translate.baidu_trans_appid318    apikey = meme_config.translate.baidu_trans_apikey319    if not appid or not apikey:320        raise MemeGeneratorException(321            "The `baidu_trans_appid` or `baidu_trans_apikey` is not set."322            "Please check your config file!"323        )324    salt = str(round(time.time() * 1000))325    sign_raw = appid + text + salt + apikey326    sign = hashlib.md5(sign_raw.encode("utf8")).hexdigest()327    params = {328        "q": text,329        "from": lang_from,330        "to": lang_to,331        "appid": appid,332        "salt": salt,333        "sign": sign,334    }335    url = "https://fanyi-api.baidu.com/api/trans/vip/translate"336    async with httpx.AsyncClient() as client:337        resp = await client.get(url, params=params)338        result = resp.json()339    return result["trans_result"][0]["dst"]340    341async def translate_microsoft(text: str, lang_from: str = "zh-CN", lang_to: str = "ja") -> str:342    if lang_to == 'jp':343        lang_to = 'ja'344    params = {345        "text": text,346        "toLang": lang_to,347    }348    url = "http://translate.ikechan8370.com/translate"349    async with httpx.AsyncClient() as client:350        resp = await client.get(url, params=params)351        result = resp.json()352    return result["translation"]["translation"]353 354def random_text() -> str:355    return random.choice(["刘一", "陈二", "张三", "李四", "王五", "赵六", "孙七", "周八", "吴九", "郑十"])356 357 358def random_image() -> BytesIO:359    text = random.choice(["😂", "😅", "🤗", "🤤", "🥵", "🥰", "😍", "😭", "😋", "😏"])360    return (361        BuildImage.new("RGBA", (500, 500), "white")362        .draw_text((0, 0, 500, 500), text, max_fontsize=400)363        .save_png()364    )365 366 367@dataclass368class TextProperties:369    fill: ColorType = "black"370    style: FontStyle = "normal"371    weight: FontWeight = "normal"372    stroke_width: int = 0373    stroke_fill: Optional[ColorType] = None374 375 376def default_template(meme: "Meme", number: int) -> str:377    return f"{number}. {'/'.join(meme.keywords)}"378 379 380def render_meme_list(381    meme_list: List[Tuple["Meme", TextProperties]],382    *,383    template: Callable[["Meme", int], str] = default_template,384    order_direction: Literal["row", "column"] = "column",385    columns: int = 4,386    column_align: Literal["left", "center", "right"] = "left",387    item_padding: Tuple[int, int] = (15, 6),388    image_padding: Tuple[int, int] = (50, 50),389    bg_color: ColorType = "white",390    fontsize: int = 30,391    fontname: str = "",392    fallback_fonts: List[str] = [],393) -> BytesIO:394    item_images: List[Text2Image] = []395    for i, (meme, properties) in enumerate(meme_list, start=1):396        text = template(meme, i)397        t2m = Text2Image.from_text(398            text,399            fontsize=fontsize,400            style=properties.style,401            weight=properties.weight,402            fill=properties.fill,403            stroke_width=properties.stroke_width,404            stroke_fill=properties.stroke_fill,405            fontname=fontname,406            fallback_fonts=fallback_fonts,407        )408        item_images.append(t2m)409    char_A = (410        Text2Image.from_text(411            "A", fontsize=fontsize, fontname=fontname, fallback_fonts=fallback_fonts412        )413        .lines[0]414        .chars[0]415    )416    num_per_col = math.ceil(len(item_images) / columns)417    column_images: List[BuildImage] = []418    for col in range(columns):419        if order_direction == "column":420            images = item_images[col * num_per_col : (col + 1) * num_per_col]421        else:422            images = [423                item_images[num * columns + col]424                for num in range((len(item_images) - col - 1) // columns + 1)425            ]426        img_w = max((t2m.width for t2m in images)) + item_padding[0] * 2427        img_h = (char_A.ascent + item_padding[1] * 2) * len(images) + char_A.descent428        image = BuildImage.new("RGB", (img_w, img_h), bg_color)429        y = item_padding[1]430        for t2m in images:431            if column_align == "left":432                x = 0433            elif column_align == "center":434                x = (img_w - t2m.width - item_padding[0] * 2) // 2435            else:436                x = img_w - t2m.width - item_padding[0] * 2437            t2m.draw_on_image(image.image, (x, y))438            y += char_A.ascent + item_padding[1] * 2439        column_images.append(image)440 441    img_w = sum((img.width for img in column_images)) + image_padding[0] * 2442    img_h = max((img.height for img in column_images)) + image_padding[1] * 2443    image = BuildImage.new("RGB", (img_w, img_h), bg_color)444    x, y = image_padding445    for img in column_images:446        image.paste(img, (x, y))447        x += img.width448    return image.save_jpg()449