Team Ai
Apppublic

horiyouta/Machine-Learning-RPG

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
main.py246 linesDownload Raw Back to root
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")