画像のタグ付けを作ってみた。
画像のタグ付けに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以上です。