見出し画像

Sanaで画像の生成を試す(Windows+CUDA、要VRAM 12GB)

※ Last update 12-21-2024
※ (1-26-2025) ComfyUI上で生成する記事を公開しています。
※ (12-21-2024) 2-1.に掲載しているコードでエラーが出たので修正しました。70行目に、image_tensorをfloat32型に変換する処理を追加しています。
※ (12-10-2024) 新たに中国語、英語、絵文字に対応した1.6Bと0.6Bのモデルが公開されているので、リンクを0-2.に追加しました。モデルを設置してコードにあるモデルのファイル名を変更し、0.6Bの場合はconfigの方も1600を600に修正してください。
※ Windows対応のフォーク版があり、こちらを利用すると1-7.~1-9.のファイル修正は不要になります(確認済み)。ただし、同様にGUIは起動できませんでした(Anacondaであれば動くかも?)。
→ https://github.com/newgenai79/Sana_win 




■ 0. 概要

▼ 0-0. はじめに

 本記事で紹介する「Sana」は10-14-2024に発表された、新しい画像生成AIです。Technical Reportの著者はNVIDIA、MIT、Tsinghua University(清華大学)に所属しています。11-21-2024にコードが公開されたので生成を試してみます。

 なお、PC上で実行しなくてもオンラインのDemoにてお試しができます。
→ https://nv-sana.mit.edu/ 

▼ 0-1. Sanaの特徴

 SanaのTechnical ReportをGoogleのNotebookLMに取り込み、対話により得た内容を簡潔に紹介します。

  • オートエンコーダの圧縮率を通常の8倍から32倍に上げ、最大4096x4096の高解像度画像を効率的に学習、生成。

  • 線形DiTを採用して計算量を削減。

  • テキストエンコーダは、従来のCLIPやT5ではなく小型LLMのGemmaを採用し、プロンプトに対する理解と推論能力を向上。

  • Flow-DPM-Solverにより推論のステップ数を削減。

  • 8bit整数量子化とTritonの利用により、処理の高速化と省リソース化。

  • Sana-0.6Bは、4K画像ではFLUX.1の100倍以上、1K解像度では40倍の高速なスループット(注:執筆時点で公開されているモデルは0.6Bではなく1.6B)。

▼ 0-2. Sanaの情報

SANA: Efficient High-Resolution Image Synthesis with Linear Diffusion Transformers



■ 1. インストール

▼ 1-1. 補足

 ハードウェア要件として「9GB VRAM is required for 0.6B model and 12GB VRAM for 1.6B model. Our later quantization version will require less than 8GB for inference.」と書かれています。現在のところ、12GBが必要のようです。

 Windowsの場合はプログラムコードの書き換えが必要なので、適宜説明します。

▼ 1-2. 本体のダウンロード

 作業ディレクトリを「\aiwork」としていますので、お好みの場所に読み替えてください。コマンドプロンプトを開き、下記のコマンドを順に実行します。

cd \aiwork
git clone https://github.com/NVlabs/Sana

▼ 1-3. 環境の構築1

 下記のコマンドを順に実行します。

cd Sana
python -m venv venv
venv\Scripts\activate
python -m pip install --upgrade pip

 続けて、下記のコマンドを実行します。ダウンロード等で若干の時間を要します。

pip install -U xformers==0.0.27.post2 --index-url https://download.pytorch.org/whl/cu121

▼ 1-4. 環境の構築2

 執筆時点では、Windows用のTritonは「pip install triton」ではインストールできません。これが原因で次の手順が詰まってしまうので、先に手動でインストールします。「pyproject.toml」の内容を見ると、「triton=3.0.0」が指定されているようです。

 下記のコマンドを実行してください。

pip install https://github.com/woct0rdho/triton-windows/releases/download/v3.1.0-windows.post5/triton-3.0.0-cp310-cp310-win_amd64.whl

 なお、執筆時点ではTritonを利用するのはGUI版のみのようです。

▼ 1-5. 環境の構築3

 続けて、下記のコマンドを実行します。ダウンロード等で若干の時間を要します。

pip install -e .

▼ 1-6. モデルのダウンロード

 Sanaのディレクトリに「checkpoints」ディレクトリを作成して、そちらへ移動してください。

 なお、512px版のモデルは下記のURLにあります。こちらを利用する場合は生成コードのcheckpoint_nameと、その次の行にあるconfigのファイル名を変更する必要があります。

▼ 1-7. ファイルの書き換え(wids_dl.py)

※ 書き換えるファイルのバックアップを推奨します。

 次に「Sana\diffusion\data\wids\wids_dl.py」のコードを、Windowsで動作するように修正します。下記のとおりに1行を変更します。

import fcntl
↓
import msvcrt

 class ULockFileの部分を、下記に差し替えます。コードはChatGPTにて作成しました。

class ULockFile:
    def __init__(self, path):
        self.path = path
        self.file = None
    def __enter__(self):
        try:
            self.file = open(self.path, "w")
            msvcrt.locking(self.file.fileno(), msvcrt.LK_LOCK, os.path.getsize(self.path) or 1)
            return self
        except Exception as e:
            if self.file:
                self.file.close()
            raise e
    def __exit__(self, exc_type, exc_val, exc_tb):
        if self.file:
            try:
                msvcrt.locking(self.file.fileno(), msvcrt.LK_UNLCK, os.path.getsize(self.path) or 1)
            finally:
                self.file.close()
                try:
                    os.unlink(self.path)
                except FileNotFoundError:
                    pass

▼ 1-8. ファイルの書き換え(wids_mmtar.py)

 次は「Sana\diffusion\data\wids\wids_mmtar.py」です。同様に処理を変更します。下記のとおりに1行を変更します。

import fcntl
↓
import msvcrt

 def keep_while_readingの部分を、下記に差し替えます。コードはChatGPTにて作成しました。

def keep_while_reading(fname, fd, phase, delay=0.0):
    assert delay == 0.0, "delay not implemented"
    if fd < 0 or fname is None:
        return
    try:
        file_size = os.path.getsize(fname) or 1
        if phase == "start":
            msvcrt.locking(fd, msvcrt.LK_LOCK, file_size)
        elif phase == "end":
            try:
                msvcrt.locking(fd, msvcrt.LK_NBLCK, file_size)
                os.unlink(fname)
            except (PermissionError, BlockingIOError):
                pass
        else:
            raise ValueError(f"Unknown phase {phase}")
    except Exception as e:
        print(f"Error occurred during file locking: {e}")

▼ 1-9. ファイルの書き換え(logger.py)

 最後に「Sana\diffusion\utils\logger.py」を変更します。下記のとおりに1行を変更します。これにより、ログファイルも出力されるようになります。

log_file = "/dev/null"
↓
log_file = "sana_log.txt"



■ 2. 画像の生成

▼ 2-1. 生成用のコード(CUI版)

 GUI版はさらに手を加えないと動作しないようなので利用を断念し、掲載されているシンプルなコードを改良します。改良にあたってはChatGPTを利用しました。

  • 生成枚数の設定を追加(デフォルトは5)。

  • そのままではコントラストが高く暗い画像が生成されるので、ヒストグラムストレッチ、コントラスト及びガンマの調整を追加(設定値は調整済み)。

  • ネガティブプロンプトを追加(効くかどうかは未確認)。

  • 保存するファイル名の変更。

  • パラメーターを画像に記録する処理を追加。

  • その他の細かい調整。

 下記のコード(sana_test.py)を「\aiwork\Sana」に保存してください。改造、再公開は自由です。

(※ 12-21-2024更新…「image_tensor.to(torch.float32)」を追加)

import os
import torch
from app.sana_pipeline import SanaPipeline
from torchvision.transforms.functional import to_pil_image
from torchvision.transforms.functional import adjust_gamma
from datetime import datetime
from PIL import PngImagePlugin

# Device configuration
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")

# Create output folder
output_dir = "output"
os.makedirs(output_dir, exist_ok=True)

# Set up the pipeline
model_directory = "checkpoints/"
checkpoint_name = "Sana_1600M_1024px.pth"
sana = SanaPipeline("configs/sana_config/1024ms/Sana_1600M_img1024.yaml")
sana.from_pretrained(f"{model_directory}{checkpoint_name}")

# Number of images to generate
num_images = 5

# Generation parameters
prompt = '''A cheerful kawaii kitten sitting under a pastel pink cherry blossom tree, with a glowing rainbow sign that sparkles with the name "Sana" in cute bubble letters. The kitten has big adorable eyes and is wearing a tiny flower crown, surrounded by floating hearts and stars'''
negative_prompt = ''
guidance_scale = 5.0 # 1-[5]-10
pag_guidance_scale = 2.0 # 1-[2]-4
num_inference_steps = 18 # 5-[18]-40

# Image generation loop
for i in range(num_images):
    # Generate a seed
    seed = torch.seed()
    generator = torch.Generator(device=device).manual_seed(seed)

    # Generate an image
    image_tensor = sana(
        prompt=prompt,
        negative_prompt=negative_prompt,
        height=1024,
        width=1024,
        guidance_scale=guidance_scale,
        pag_guidance_scale=pag_guidance_scale,
        num_inference_steps=num_inference_steps,
        generator=generator,
    )

    # Histogram stretching
    #min_val = image_tensor.min()
    min_val = -1.51
    #max_val = image_tensor.max()
    max_val = 1.51
    #print(f"min_val={min_val} max_val={max_val} ")
    stretched = (image_tensor - min_val) / (max_val - min_val + 1e-5)

    # Contrast enhancement
    contrast_factor = 1.5 # n < [1.0] < m
    adjusted = 0.5 + contrast_factor * (stretched - 0.5)
    # Re-clamp pixel values
    image_tensor = adjusted.clamp(0, 1)

    # Gamma correction
    gamma = 0.8 # 0.2 < [1.0] < n
    image_tensor = image_tensor.clamp(0, 1)
    image_tensor = adjust_gamma(image_tensor, gamma)

    # Convert Tensor to Pillow image
    image_tensor = image_tensor.to(torch.float32)
    image = to_pil_image(image_tensor.squeeze().clamp(0, 1))

    # Get current timestamp and create a filename
    timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
    filename = f"{output_dir}/{timestamp}-{seed}-sana.png"

    # Set metadata
    img_parameters = (
        f"{prompt} \n"
        f"Negative prompt: {negative_prompt} \n"
        f"Steps: {num_inference_steps}, "
        f"Guidance Scale: {guidance_scale}, "
        f"PAG Guidance Scale: {pag_guidance_scale}, "
        f"Seed: {seed}, "
        f"Model: {checkpoint_name}"
    )
    metadata = PngImagePlugin.PngInfo()
    metadata.add_text("parameters", img_parameters)

    # Save an image
    image.save(filename, "PNG", pnginfo=metadata)
    print(f"Saved: {filename} with metadata.")

▼ 2-2. 生成の実行

 コマンドプロンプトを新規で開き、下記のコマンドを順に実行してください。1.の手順で使用したウインドウが残っている場合、これを実行する必要はありません。

cd \aiwork\Sana
venv\Scripts\activate

 必要に応じて「sana_test.py」の内容を書き換えてから、下記のコマンドを実行します。

python sana_test.py

 初回の実行時は必要なファイルをダウンロードするため、少々時間を要します。準備が終わると生成を開始し、「\aiwork\Sana\output」に保存します。下記の画像は生成例です。

A cheerful kawaii kitten sitting under a pastel pink cherry blossom tree, with a glowing rainbow sign that sparkles with the name "Sana" in cute bubble letters. The kitten has big adorable eyes and is wearing a tiny flower crown, surrounded by floating hearts and stars
Steps: 18, Guidance Scale: 5.0, PAG Guidance Scale: 2.0, Seed: 68930337598400, Model: Sana_1600M_1024px.pth

▼ 2-3. VRAMの使用量

 VRAMが12GBの場合はギリギリで、不要なプログラムを閉じても共有GPUメモリに少しはみ出るかもしれません。

実行中のGPUメモリの使用量

▼ 2-4. 設定の変更等について

 基本的に、「prompt」以外を書き換える必要はありません。記号が使えるようにトリプルクォートを用いているので、「'''」と「'''」の間にプロンプトを記述してください。

 デフォルトでは画像を5枚生成するようになっているので、好みで「num_images」の値を変更してください。

 生成パラメーターの「guidance_scale」「pag_guidance_scale」「num_inference_step」は好みで変更できます。ただし、通常はデフォルト値のままで良いと思います。「negative_prompt」が有効であるかどうかは未確認です(利用しているdiffusersのバージョンに依存?)。

 生成後の処理として「Histogram stretching」「Contrast enhancement」「Gamma correction」を追加して、筆者が値を調整しました。これらの処理を行わないと、コントラストが高く暗い画像になってしまいます。ガンマ値は低くするほど明るくなるので注意してください。



■ 3. おまけ

▼ 3-1. Sanaのプロンプト

 基本的には自然言語で良いと思います。GUI版のコードを見ると、プロンプトの書き方について若干のヒントが得られます。

Sana/app/app_sana.py
https://github.com/NVlabs/Sana/blob/main/app/app_sana.py 

 一つは47行目あたりに「style_list」があり、様々なスタイルの基本的なpromptとnegative_promptが記されているので参考になりそうです。

 もう一つは319行目あたりに「examples」があり、Demoでプリセットされているプロンプトが書いてあります。

▼ 3-2. おまけ画像

 生成した画像とパラメーターを掲載します。

A fairy girl with translucent wings gives a playful wink and blows a kiss in an enchanted forest glade. Her tilted flower crown catches the morning light as fireflies dance around her in the pastel-hued scene. Delicate wildflowers and glowing pollen float in the shimmering air.
Steps: 18, Guidance Scale: 5.0, PAG Guidance Scale: 2.0, Seed: 74641912530700, Model: Sana_1600M_1024px.pth
A manga-style pre-teen girl with light brown pigtails and gray eyes snuggles close to a large golden retriever. She wears a light green plaid dress with frills and a chest ribbon, paired with brown boots. The scene is set on a grassy hillside overlooking a distant lake, with scattered clouds in the blue sky and colorful flowers around them. Drawn in a vibrant pastel palette from a low angle, close-up perspective.
Steps: 18, Guidance Scale: 5.0, PAG Guidance Scale: 2.0, Seed: 76935840839700, Model: Sana_1600M_1024px.pth
A buff carrot superhero with muscular legs flexes in a garden, showing off to a laid-back tomato who lounges against a fence wearing sunglasses. Nearby, a clumsy potato stumbles around with tiny gardening tools, while a goofy cucumber attempts gymnastics on a trellis. The veggies have cartoonish faces with exaggerated expressions, and sparkles of afternoon sun make their antics even sillier.
Steps: 18, Guidance Scale: 5.0, PAG Guidance Scale: 2.0, Seed: 79671462399100, Model: Sana_1600M_1024px.pth

▼ 3-3. おまけ画像2

 サンプルコードやDemoに掲載されているプロンプトを利用しました。

a cyberpunk cat with a neon sign that says "Sana"
Steps: 18, Guidance Scale: 5.0, PAG Guidance Scale: 2.0, Seed: 80658416527100, Model: Sana_1600M_1024px.pth
👧 with 🌹 in the ❄️
Steps: 18, Guidance Scale: 5.0, PAG Guidance Scale: 2.0, Seed: 81072049976900, Model: Sana_1600M_1024px.pth
a stunning and luxurious bedroom carved into a rocky mountainside seamlessly blending nature with modern design with a plush earth-toned bed textured stone walls circular fireplace massive uniquely shaped window framing snow-capped mountains dense forests
Steps: 18, Guidance Scale: 5.0, PAG Guidance Scale: 2.0, Seed: 81431696380000, Model: Sana_1600M_1024px.pth



■ 4. その他

 私が書いた他の記事は、メニューよりたどってください。

 noteのアカウントはメインの@Mayu_Hiraizumiに紐付けていますが、記事に関することはサブアカウントの@riddi0908までお願いします。

いいなと思ったら応援しよう!