メインコンテンツへスキップ

画像のタグ付けを作ってみた。

    画像のタグ付けにonnxモデルを使用したTaggerという機能があるが、
    実際に実装してみました。

    index.html

    <!DOCTYPE html>
    <html lang="ja">
    <head>
        <meta charset="UTF-8">
        <meta name="viewport" content="width=device-width, initial-scale=1.0">
        <title>画像タグ生成</title>
        <link rel="icon" href="data:,">
        <script src="https://cdn.jsdelivr.net/npm/onnxruntime-web/dist/ort.min.js"></script>
        <link rel="stylesheet" href="styles.css">
    </head>
    <body>
        <div class="header">
            <h1>jtagger - 画像タグ生成</h1>
        </div>
        
        <div class="input-section">
            <div class="file-input-container">
                <input type="file" id="imageInput" accept="image/*" class="file-input">
                <label for="imageInput" class="file-label">画像を選択</label>
            </div>
        </div>
    
        <div id="thumbnailContainer" style="display: none; margin: 10px 0;"></div>
    
        <div class="threshold-section">
            <label for="thresholdSlider">一般タグしきい値: <span id="thresholdValue">0.35</span></label>
            <input type="range" id="thresholdSlider" min="0" max="1" step="0.05" value="0.35">
            <div class="threshold-controls">
                <button id="thresholdDecrement">-</button>
                <button id="thresholdIncrement">+</button>
            </div>
        </div>
    
        <div class="threshold-section">
            <label for="characterThresholdSlider">キャラクタータグしきい値: <span id="characterThresholdValue">0.35</span></label>
            <input type="range" id="characterThresholdSlider" min="0" max="1" step="0.05" value="0.35">
            <div class="threshold-controls">
                <button id="characterThresholdDecrement">-</button>
                <button id="characterThresholdIncrement">+</button>
            </div>
        </div>
    
        <div class="button-section">
            <button id="generateTagsButton">タグ生成</button>
        </div>
    
        <div class="progress-section">
            <div id="initializationMessage">初期化中...</div>
            <div id="progress" style="display: none;">
                <progress id="progressBar" value="0" max="100"></progress>
                <p id="progressMessage">タグ生成中...</p>
            </div>
        </div>
    
        <div class="results-section">
            <textarea id="tagResults"></textarea>
            <div class="copy-section">
                <button id="copyButton">タグをコピー</button>
            </div>
        </div>
    
        <script src="ui.js"></script>
        <script>
            // ページ読み込み時の初期化
            window.addEventListener('load', async () => {
                try {
                    await initialize();
                    onInitializationComplete();
                } catch (error) {
                    onInitializationError(error);
                }
            });
        </script>
        <script src="jtagger.js"></script>
    </body>
    </html> 

    styles.css

    .file-input-container {
        display: flex;
        align-items: center;
        gap: 1rem;
        margin-bottom: 1rem;
        position: relative;
    }
    
    .file-input {
        display: none;
    }
    
    .file-label {
        padding: 0.5rem 1rem;
        background-color: #4CAF50;
        color: white;
        border-radius: 4px;
        cursor: pointer;
        transition: background-color 0.3s;
    }
    
    .file-label:hover {
        background-color: #45a049;
    }
    
    .thumbnail-container {
        width: 200px;
        height: 200px;
        border: 1px solid #ccc;
        border-radius: 4px;
        overflow: hidden;
        display: none;
        flex-shrink: 0;
        background-color: #f5f5f5;
        margin: 1rem 0;
    }
    
    .thumbnail-container canvas {
        width: 100%;
        height: 100%;
        display: block;
    }
    
    .thumbnail-container img {
        width: 100%;
        height: 100%;
        object-fit: contain;
        display: block;
    }
    
    .input-section {
        margin-bottom: 1rem;
    }
    
    #tagResults {
        width: 100%;
        height: 10em;  /* 10行分の高さ */
        padding: 0.5rem;
        margin-bottom: 1rem;
        border: 1px solid #ccc;
        border-radius: 4px;
        resize: none;  /* リサイズを無効化 */
        font-family: inherit;
        font-size: inherit;
        line-height: 1.5;
    }
    
    .threshold-section {
        display: flex;
        align-items: center;
        gap: 1rem;
        margin-bottom: 1rem;
    }
    
    .threshold-controls {
        display: flex;
        gap: 0.5rem;
        margin-left: 0;  /* 右寄せを解除 */
    }
    
    .threshold-controls button {
        padding: 0.25rem 0.5rem;
        background-color: #4CAF50;
        color: white;
        border: none;
        border-radius: 4px;
        cursor: pointer;
        transition: background-color 0.3s;
    }
    
    .threshold-controls button:hover {
        background-color: #45a049;
    } 

    ui.js

    // UI関連の関数
    
    // ボタンの有効/無効化
    function enableButtons() {
        document.getElementById('imageInput').disabled = false;
        document.getElementById('generateTagsButton').disabled = false;
        document.getElementById('thresholdSlider').disabled = false;
        document.getElementById('thresholdDecrement').disabled = false;
        document.getElementById('thresholdIncrement').disabled = false;
        document.getElementById('characterThresholdSlider').disabled = false;
        document.getElementById('characterThresholdDecrement').disabled = false;
        document.getElementById('characterThresholdIncrement').disabled = false;
        document.getElementById('copyButton').disabled = false;
    }
    
    function disableButtons() {
        document.getElementById('imageInput').disabled = true;
        document.getElementById('generateTagsButton').disabled = true;
        document.getElementById('thresholdSlider').disabled = true;
        document.getElementById('thresholdDecrement').disabled = true;
        document.getElementById('thresholdIncrement').disabled = true;
        document.getElementById('characterThresholdSlider').disabled = true;
        document.getElementById('characterThresholdDecrement').disabled = true;
        document.getElementById('characterThresholdIncrement').disabled = true;
        document.getElementById('copyButton').disabled = true;
    }
    
    // プログレスバーの表示/非表示
    function showProgress() {
        const progress = document.getElementById('progress');
        progress.style.display = 'block';
        updateProgressBar(0);
    }
    
    function hideProgress() {
        const progress = document.getElementById('progress');
        progress.style.display = 'none';
    }
    
    // プログレスバーの更新
    function updateProgressBar(value) {
        const progressBar = document.getElementById('progressBar');
        progressBar.value = value;
    }
    
    // プログレスメッセージの更新
    function updateProgressMessage(message) {
        const progressMessage = document.getElementById('progressMessage');
        progressMessage.textContent = message;
    }
    
    // タグ結果の表示
    function displayTags(tags) {
        const tagResults = document.getElementById('tagResults');
        const tagNames = tags.map(tag => tag.tag.split(',')[1]); // 2列目のタグ名のみを抽出
        tagResults.value = tagNames.join(', '); // タグ名をカンマ区切りで表示
    }
    
    // タグのコピー
    async function copyTags() {
        const tagResults = document.getElementById('tagResults');
        const copyButton = document.getElementById('copyButton');
        
        try {
            await navigator.clipboard.writeText(tagResults.value);
            const originalText = copyButton.textContent;
            copyButton.textContent = 'コピーしました!';
            setTimeout(() => {
                copyButton.textContent = originalText;
            }, 2000);
        } catch (error) {
            console.error('コピーに失敗しました:', error);
            copyButton.textContent = 'コピーに失敗しました';
            setTimeout(() => {
                copyButton.textContent = 'タグをコピー';
            }, 2000);
        }
    }
    
    // 一般タグしきい値の更新
    function updateThresholdValue(value) {
        const thresholdValue = document.getElementById('thresholdValue');
        thresholdValue.textContent = value.toFixed(2);
    }
    
    // キャラクタータグしきい値の更新
    function updateCharacterThresholdValue(value) {
        const characterThresholdValue = document.getElementById('characterThresholdValue');
        characterThresholdValue.textContent = value.toFixed(2);
    }
    
    // 初期化メッセージの更新
    function updateInitializationMessage(message) {
        const initializationMessage = document.getElementById('initializationMessage');
        initializationMessage.textContent = message;
    }
    
    // 初期化完了時の処理
    function onInitializationComplete() {
        document.getElementById('initializationMessage').style.display = 'none';
        enableButtons();
    }
    
    // 初期化エラー時の処理
    function onInitializationError(error) {
        console.error('初期化エラー:', error);
        updateInitializationMessage('初期化に失敗しました。ページを再読み込みしてください。');
    } 

    jtagger.js

    // 初期値を設定
    const defaultThreshold = 0.35;
    const defaultCharacterThreshold = 0.35;
    
    // モデルとタグリストの読み込み
    async function loadModel() {
        console.log('モデルとタグリストを読み込み中...');
        try {
            // ONNXセッションの初期化
            if (!window.onnxSession) {
                console.log('ONNXセッションを初期化中...');
                window.onnxSession = await ort.InferenceSession.create('model.onnx');
                console.log('ONNXセッション初期化完了');
            }
    
            // タグリストの読み込み
            if (!window.tagList) {
                console.log('タグリストを読み込み中...');
                window.tagList = await loadTagList('selected_tags.csv');
                console.log('タグリスト読み込み完了');
            }
    
            return true;
        } catch (error) {
            console.error('モデルとタグリストの読み込みに失敗:', error);
            return false;
        }
    }
    
    // 初期化処理
    async function initialize() {
        console.log('initialize() called');
        try {
            // スライダーの初期値を設定
            const thresholdSlider = document.getElementById('thresholdSlider');
            const characterThresholdSlider = document.getElementById('characterThresholdSlider');
            
            thresholdSlider.value = defaultThreshold;
            updateThresholdValue(defaultThreshold);
            characterThresholdSlider.value = defaultCharacterThreshold;
            updateCharacterThresholdValue(defaultCharacterThreshold);
    
            console.log('jtagger 初期化完了');
            return true;
        } catch (error) {
            console.error('jtagger 初期化失敗:', error);
            return false;
        }
    }
    
    // ファイルが存在しない場合にダウンロード
    async function downloadFileIfNotExists(filename, url) {
        try {
            await fetch(filename, { method: 'HEAD' });
            console.log(`${filename} が存在します。`);
            return true; // ファイルが存在する場合trueを返す
        } catch (error) {
            console.log(`${filename} をダウンロードします。`);
            try{
                const response = await fetch(url);
                const blob = await response.blob();
                const a = document.createElement('a');
                a.href = window.URL.createObjectURL(blob);
                a.download = filename;
                a.style.display = 'none';
                document.body.appendChild(a);
                a.click();
                document.body.removeChild(a);
                console.log(`${filename} のダウンロードが完了しました。`);
                return false; // ファイルをダウンロードした場合falseを返す
            } catch(error){
                console.error(`${filename} のダウンロードに失敗しました。`, error);
                return null; // ダウンロード失敗時はnullを返す
            }
        }
    }
    
    // タグリストの読み込み
    async function loadTagList(filePath) {
        const response = await fetch(filePath);
        const text = await response.text();
        return text.split('\n');
    }
    
    // 画像のタグ生成
    async function generateImageTags() {
        try {
            // モデルの初期化
            updateProgressMessage('モデルを初期化中...');
            updateProgressBar(0);
            const modelLoaded = await loadModel();
            if (!modelLoaded) {
                throw new Error('モデルの初期化に失敗しました');
            }
            updateProgressBar(20);
    
            // ボタンを無効化
            disableButtons();
            
            // 進捗表示を開始
            showProgress();
            updateProgressMessage('画像を読み込み中...');
    
            // コンソールにメッセージを表示
            console.log('タグ生成を開始します。');
    
            let modelExists = await downloadFileIfNotExists('model.onnx', 'https://huggingface.co/SmilingWolf/wd-eva02-large-tagger-v3/resolve/main/model.onnx');
            let tagsExists = await downloadFileIfNotExists('selected_tags.csv', 'https://huggingface.co/SmilingWolf/wd-eva02-large-tagger-v3/resolve/main/selected_tags.csv');
    
            if (modelExists === false || tagsExists === false) {
                updateProgressMessage('モデルを再初期化中...');
                updateProgressBar(30);
                const modelLoaded = await loadModel();
                if (!modelLoaded) {
                    throw new Error('モデルの再初期化に失敗しました');
                }
                updateProgressBar(40);
            }
    
            const imageInput = document.getElementById('imageInput');
            const file = imageInput.files[0];
            if (!file) {
                throw new Error('画像が選択されていません');
            }
    
            // 画像のプリプロセス
            updateProgressMessage('画像を処理中...');
            const imageData = await preprocessImage(file, (progress) => {
                requestAnimationFrame(() => {
                    updateProgressBar(40 + progress * 20);
                    updateProgressMessage(`画像を処理中... ${Math.round(progress * 100)}%`);
                });
            });
    
            // 推論処理
            updateProgressMessage('タグを生成中...');
            const tags = await runInference(imageData, (progress) => {
                requestAnimationFrame(() => {
                    updateProgressBar(60 + progress * 30);
                    updateProgressMessage(`タグを生成中... ${Math.round(progress * 100)}%`);
                });
            });
    
            // タグの表示
            updateProgressMessage('タグを表示中...');
            displayTags(tags);
            updateProgressBar(100);
    
            // 進捗表示を終了
            setTimeout(() => {
                hideProgress();
            }, 1000);
    
            // ボタンを有効化
            enableButtons();
    
            // コンソールにメッセージを表示
            console.log('タグ生成が完了しました。');
    
        } catch (error) {
            console.error('タグ生成失敗:', error);
            updateProgressMessage(`エラーが発生しました: ${error.message}`);
            hideProgress();
            enableButtons();
        }
    }
    
    // 画像の前処理
    async function preprocessImage(file, progressCallback) {
        return new Promise((resolve, reject) => {
            const img = new Image();
            img.onload = async () => {
                try {
                    const canvas = document.createElement('canvas');
                    canvas.width = 448;
                    canvas.height = 448;
                    const ctx = canvas.getContext('2d');
                    ctx.drawImage(img, 0, 0, 448, 448);
                    const imageData = ctx.getImageData(0, 0, 448, 448);
    
                    // RGBAデータをRGBデータに変換
                    const rgbData = new Uint8Array(448 * 448 * 3);
                    for (let i = 0, j = 0; i < imageData.data.length; i += 4, j += 3) {
                        rgbData[j] = imageData.data[i];
                        rgbData[j + 1] = imageData.data[i + 1];
                        rgbData[j + 2] = imageData.data[i + 2];
                        if (progressCallback && i % 1000 === 0) {
                            await new Promise(resolve => setTimeout(resolve, 0));
                            progressCallback(i / imageData.data.length);
                        }
                    }
    
                    resolve({ data: rgbData, width: 448, height: 448 });
                } catch (error) {
                    reject(error);
                }
            };
            img.onerror = () => reject(new Error('画像の読み込みに失敗しました'));
            img.src = URL.createObjectURL(file);
        });
    }
    
    // 推論の実行
    async function runInference(imageData, progressCallback) {
        try {
            const input = new ort.Tensor('float32', new Float32Array(imageData.data), [1, imageData.height, imageData.width, 3]);
            const outputMap = await window.onnxSession.run({ 'input': input });
            const output = outputMap['output'].data;
    
            const tags = [];
            const thresholdSlider = parseFloat(document.getElementById('thresholdSlider').value);
            const characterThresholdSlider = parseFloat(document.getElementById('characterThresholdSlider').value);
    
            for (let i = 0; i < output.length; i++) {
                const tag = window.tagList[i];
                const isCharacterTag = tag.startsWith('1girl') || tag.startsWith('1boy') || tag.startsWith('2girls') || tag.startsWith('2boys');
                const threshold = isCharacterTag ? characterThresholdSlider : thresholdSlider;
    
                if (output[i] > threshold) {
                    tags.push({ tag: tag, confidence: output[i], isCharacterTag: isCharacterTag });
                    if (progressCallback && i % 100 === 0) {
                        await new Promise(resolve => setTimeout(resolve, 0));
                        progressCallback(i / output.length);
                    }
                }
            }
            return tags;
        } catch (error) {
            throw new Error(`推論処理中にエラーが発生しました: ${error.message}`);
        }
    }
    
    // タグの表示
    function displayTags(tags) {
        const resultTextbox = document.getElementById('tagResults');
        const tagNames = tags.map(tag => tag.tag.split(',')[1]); // 2列目のタグ名のみを抽出
        resultTextbox.value = tagNames.join(', '); // タグ名をカンマ区切りで表示
    }
    
    // 進捗表示を開始
    function showProgress() {
        document.getElementById('progress').style.display = 'block';
    }
    
    // 進捗表示を終了
    function hideProgress() {
        document.getElementById('progress').style.display = 'none';
    }
    
    // プログレスバーの更新
    function updateProgressBar(value) {
        document.getElementById('progressBar').value = value;
    }
    
    // ボタンを無効化
    function disableButtons() {
        document.getElementById('imageInput').disabled = true;
        document.getElementById('generateTagsButton').disabled = true;
    }
    
    // ボタンを有効化
    function enableButtons() {
        document.getElementById('imageInput').disabled = false;
        document.getElementById('generateTagsButton').disabled = false;
    }
    
    // スライダーの値を表示
    const thresholdSlider = document.getElementById('thresholdSlider');
    const thresholdValue = document.getElementById('thresholdValue');
    
    thresholdSlider.addEventListener('input', function () {
        thresholdValue.textContent = this.value;
    });
    
    // スピンボタンの処理
    thresholdSlider.addEventListener('change', function () {
        thresholdValue.textContent = this.value;
    });
    
    // キャラクタータグスライダーの値を表示
    const characterThresholdSlider = document.getElementById('characterThresholdSlider');
    const characterThresholdValue = document.getElementById('characterThresholdValue');
    
    characterThresholdSlider.addEventListener('input', function () {
        characterThresholdValue.textContent = this.value;
    });
    
    // キャラクタータグスピンボタンの処理
    characterThresholdSlider.addEventListener('change', function () {
        characterThresholdValue.textContent = this.value;
    });
    
    // イベントリスナーの設定
    window.addEventListener('load', async () => {
        console.log('loadイベントが発生しました');
        
        // 要素の取得を確実に行う
        const imageInput = document.querySelector('#imageInput');
        const thumbnailContainer = document.querySelector('#thumbnailContainer');
        const generateTagsButton = document.querySelector('#generateTagsButton');
        const copyButton = document.querySelector('#copyButton');
    
        console.log('要素の取得結果:');
        console.log('imageInput:', imageInput);
        console.log('thumbnailContainer:', thumbnailContainer);
        console.log('generateTagsButton:', generateTagsButton);
        console.log('copyButton:', copyButton);
    
        // 生成ボタンのイベントリスナー
        if (generateTagsButton) {
            generateTagsButton.addEventListener('click', generateImageTags);
        }
    
        // コピーボタンのイベントリスナー
        if (copyButton) {
            copyButton.addEventListener('click', copyTags);
        }
    
        // 画像選択時のサムネイル表示
        if (imageInput && thumbnailContainer) {
            console.log('イベントリスナーを設定します');
            imageInput.addEventListener('change', function(e) {
                console.log('画像選択イベントが発生しました');
                const file = e.target.files[0];
                console.log('選択されたファイル:', file);
                
                if (file) {
                    const img = new Image();
                    img.onload = function() {
                        console.log('画像を読み込みました');
                        const canvas = document.createElement('canvas');
                        canvas.width = 200;
                        canvas.height = 200;
                        const ctx = canvas.getContext('2d');
                        
                        // アスペクト比を保持して描画
                        const scale = Math.min(200 / img.width, 200 / img.height);
                        const x = (200 - img.width * scale) / 2;
                        const y = (200 - img.height * scale) / 2;
                        
                        ctx.drawImage(img, x, y, img.width * scale, img.height * scale);
                        
                        // サムネイルを表示
                        thumbnailContainer.innerHTML = '';
                        thumbnailContainer.appendChild(canvas);
                        thumbnailContainer.style.display = 'block';
                        console.log('サムネイルを表示しました');
                    };
                    img.src = URL.createObjectURL(file);
                } else {
                    console.log('ファイルが選択されていません');
                    thumbnailContainer.style.display = 'none';
                    thumbnailContainer.innerHTML = '';
                }
            });
        } else {
            console.log('必要な要素が見つかりません');
            console.log('HTMLの構造を確認してください:');
            console.log(document.body.innerHTML);
        }
    
        // 初期化処理
        try {
            await initialize();
            onInitializationComplete();
        } catch (error) {
            onInitializationError(error);
        }
    });
    
    

    start_server.bat

    @echo off
    chcp 65001 > nul
    echo UTF-8 モードに変更しました。
    echo ローカルサーバーを起動します...
    start http://localhost:8080/index.html
    npx http-server -c-1
    pause

    以上です。

    あなたへのおすすめ