horiyouta/Machine-Learning-RPG
0
1from fastapi import FastAPI, Body2from fastapi.staticfiles import StaticFiles3from typing import List, Dict, Any4import torch5import torch.nn as nn6import torch.optim as optim7from torchvision import transforms8from torch.utils.data import DataLoader9import numpy as np10import base6411from io import BytesIO12from PIL import Image13import random14from datasets import load_dataset15 16# FastAPIアプリケーションインスタンスを作成17app = FastAPI()18 19# --- 動的なプレイヤーモデル (変更なし) ---20class PlayerModel(nn.Module):21 def __init__(self, layer_configs):22 super(PlayerModel, self).__init__()23 self.layers = nn.ModuleList()24 self.architecture_info = []25 self.hookable_layers = {}26 27 in_channels = 128 feature_map_size = 2829 is_flattened = False30 31 for i, config in enumerate(layer_configs):32 layer_type = config['type']33 name = f"{layer_type.lower()}_{len([info for info in self.architecture_info if info['type'] == layer_type])}"34 35 if layer_type in ['Conv2d', 'MaxPool2d', 'AvgPool2d']:36 is_flattened = False37 if layer_type == 'Conv2d':38 out_channels = config['params']['out_channels']39 kernel_size = config['params']['kernel_size']40 layer = nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, padding=kernel_size//2)41 self.layers.append(layer)42 self.hookable_layers[name] = layer43 in_channels = out_channels44 self.architecture_info.append({"type": "Conv2d", "name": name, "shape": [out_channels, feature_map_size, feature_map_size]})45 else:46 kernel_size = config['params']['kernel_size']47 if layer_type == 'MaxPool2d':48 layer = nn.MaxPool2d(kernel_size=kernel_size, stride=kernel_size)49 else:50 layer = nn.AvgPool2d(kernel_size=kernel_size, stride=kernel_size)51 self.layers.append(layer)52 self.hookable_layers[name] = layer53 feature_map_size //= kernel_size54 self.architecture_info.append({"type": layer_type, "name": name, "shape": [in_channels, feature_map_size, feature_map_size]})55 elif layer_type in ['ReLU', 'Dropout']:56 if layer_type == 'ReLU':57 self.layers.append(nn.ReLU())58 else:59 p = config['params']['p']60 self.layers.append(nn.Dropout(p=p))61 self.architecture_info.append({"type": layer_type, "name": name})62 elif layer_type == 'Flatten':63 if not is_flattened:64 layer = nn.Flatten()65 self.layers.append(layer)66 self.hookable_layers[name] = layer67 flat_features = in_channels * feature_map_size * feature_map_size68 in_channels = flat_features69 self.architecture_info.append({"type": "Flatten", "name": name, "shape": [flat_features]})70 is_flattened = True71 elif layer_type in ['Linear', 'ResidualBlock']:72 if not is_flattened:73 auto_flatten_name = f"auto_flatten_{i}"74 self.layers.append(nn.Flatten())75 flat_features = in_channels * feature_map_size * feature_map_size76 in_channels = flat_features77 self.architecture_info.append({"type": "Flatten", "name": auto_flatten_name, "shape": [flat_features]})78 is_flattened = True79 if layer_type == 'Linear':80 out_features = config['params']['out_features']81 layer = nn.Linear(in_channels, out_features)82 in_channels = out_features83 else:84 features = in_channels85 layer = nn.Linear(features, features)86 self.layers.append(layer)87 self.hookable_layers[name] = layer88 self.architecture_info.append({"type": layer_type, "name": name, "shape": [in_channels]})89 90 if not self.layers or not isinstance(self.layers[-1], nn.Linear) or self.layers[-1].out_features != 10:91 if not is_flattened:92 self.layers.append(nn.Flatten())93 final_in_features = in_channels * feature_map_size * feature_map_size94 else:95 final_in_features = in_channels96 output_layer = nn.Linear(final_in_features, 10)97 self.layers.append(output_layer)98 self.hookable_layers["linear_output"] = output_layer99 self.architecture_info.append({"type": "Linear", "name": "linear_output", "shape": [10]})100 101 def forward(self, x):102 for layer in self.layers:103 x = layer(x)104 return x105 106# --- グローバル変数とデータ準備 (ステートレス対応) ---107# これらの変数はサーバー起動時に一度だけ初期化され、リクエスト間で変更されない定数として扱う108device = torch.device("cpu")109mnist_dataset = load_dataset("mnist")110transform = transforms.Compose([transforms.ToTensor()])111 112def apply_transforms(examples):113 examples['image'] = [transform(image.convert("L")) for image in examples['image']]114 return examples115 116mnist_dataset.set_transform(apply_transforms)117train_subset = mnist_dataset['train'].select(range(1000))118train_loader = DataLoader(train_subset, batch_size=32, shuffle=True)119 120test_images = []121test_subset_for_inference = mnist_dataset['test'].shuffle().select(range(1000))122for item in test_subset_for_inference:123 image_tensor = item['image'].unsqueeze(0)124 label_tensor = torch.tensor(item['label'])125 test_images.append((image_tensor, label_tensor))126 127# --- バックエンドロジック (ステートレス関数) ---128 129def get_enemy():130 """新しい敵の画像(base64)と正解ラベルを返す。サーバー側では状態を保持しない。"""131 image_tensor, label_tensor = random.choice(test_images)132 133 img_pil = transforms.ToPILImage()(image_tensor.squeeze(0))134 buffered = BytesIO()135 img_pil.save(buffered, format="PNG")136 img_str = base64.b64encode(buffered.getvalue()).decode("utf-8")137 138 return {139 "image_b64": "data:image/png;base64," + img_str,140 "label": label_tensor.item()141 }142 143def run_inference(layer_configs: list, enemy_image_b64: str, enemy_label: int):144 """145 リクエストごとにモデルを構築・訓練し、与えられた敵データで推論を実行する。146 サーバー側では状態を一切保持しない。147 """148 # 1. モデルをその場で構築し、訓練する149 if not layer_configs:150 return {"error": "モデルが空です。"}151 try:152 model = PlayerModel(layer_configs).to(device)153 optimizer = optim.Adam(model.parameters(), lr=0.001)154 loss_fn = nn.CrossEntropyLoss()155 156 model.train()157 for epoch in range(3): # 毎回3エポック学習158 for batch in train_loader:159 data, target = batch['image'].to(device), batch['label'].to(device)160 optimizer.zero_grad()161 output = model(data)162 loss = loss_fn(output, target)163 loss.backward()164 optimizer.step()165 print("On-the-fly training for inference finished.")166 except Exception as e:167 print(f"Error during on-the-fly training: {e}")168 return {"error": f"推論中のモデル構築・訓練エラー: {e}"}169 170 # 2. クライアントから送られてきた敵画像で推論する171 model.eval()172 173 # Base64文字列から画像テンソルにデコード174 try:175 header, encoded = enemy_image_b64.split(",", 1)176 image_data = base64.b64decode(encoded)177 image_pil = Image.open(BytesIO(image_data)).convert("L")178 image_tensor = transforms.ToTensor()(image_pil).unsqueeze(0).to(device)179 except Exception as e:180 print(f"Error decoding enemy image: {e}")181 return {"error": f"敵画像のデコードエラー: {e}"}182 183 # 3. 推論と中間出力のキャプチャ184 intermediate_outputs = {}185 hooks = []186 def get_hook(name):187 def hook(model, input, output):188 intermediate_outputs[name] = output.detach().cpu().clone().numpy().tolist()189 return hook190 191 for name, layer in model.hookable_layers.items():192 hooks.append(layer.register_forward_hook(get_hook(name)))193 194 with torch.no_grad():195 output = model(image_tensor)196 197 for h in hooks: h.remove()198 199 probabilities = torch.nn.functional.softmax(output, dim=1)200 prediction = torch.argmax(probabilities, dim=1).item()201 confidence = probabilities[0, prediction].item()202 203 intermediate_outputs['input'] = image_tensor.cpu().numpy().tolist()204 205 weights = {}206 for name, layer in model.hookable_layers.items():207 if isinstance(layer, (nn.Linear, nn.Conv2d)):208 if hasattr(layer, 'weight') and hasattr(layer, 'bias'):209 weights[name + '_w'] = layer.weight.cpu().detach().numpy().tolist()210 weights[name + '_b'] = layer.bias.cpu().detach().numpy().tolist()211 212 is_correct = (prediction == enemy_label)213 214 # 4. 結果をクライアントに返す215 return {216 "prediction": prediction, 217 "label": enemy_label, 218 "is_correct": is_correct,219 "confidence": confidence,220 "image_b64": enemy_image_b64, # 受け取った画像をそのまま返す221 "architecture": [{"type": "Input", "name": "input", "shape": [1, 28, 28]}] + model.architecture_info,222 "outputs": intermediate_outputs,223 "weights": weights224 }225 226# --- FastAPI Endpoints ---227@app.get("/api/get_enemy")228async def get_enemy_endpoint():229 return get_enemy()230 231@app.post("/api/run_inference")232async def run_inference_endpoint(payload: Dict[str, Any] = Body(...)):233 """234 クライアントからモデル構成と敵データを受け取り、推論結果を返すエンドポイント。235 """236 layer_configs = payload.get("layer_configs")237 enemy_image_b64 = payload.get("enemy_image_b64")238 enemy_label = payload.get("enemy_label")239 240 if not all([layer_configs, enemy_image_b64, enemy_label is not None]):241 return {"error": "リクエストのパラメータが不足しています。"}242 243 return run_inference(layer_configs, enemy_image_b64, enemy_label)244 245# --- 静的ファイルの配信 ---246app.mount("/", StaticFiles(directory="web", html=True), name="static")