#!/usr/bin/env python3
"""آموزش کامل مدل دسته‌بندی زنبیل — یک اجرا، بدون دخالت، قابل ازسرگیری.

نسخه اسکریپتیِ نوت‌بوک، با سه تفاوت که روی یک GPU واقعی مهم‌اند:

* تصاویر **موازی** رمزگشایی می‌شوند. روی T4 گلوگاه خود GPU بود، پس تک‌رشته‌ای
  اهمیتی نداشت؛ روی کارت سریع‌تر ورق برمی‌گردد و کارت بیکار منتظر CPU می‌ماند.
  اندازه‌گیری: هر تصویر ~۳ms کار CPU، یعنی تک‌رشته‌ای سقف ~۳۳۰ تصویر/ثانیه.
* اگر VRAM جا داشته باشد، بردارهای آموزش روی خود کارت می‌مانند به‌جای انتقال
  بچ‌به‌بچ.
* هر مرحله چک‌پوینت دارد. قطع برق یا ریبوت یعنی از همان‌جا ادامه، نه از صفر.

    python3 train_local.py --data /data/zambil --shards http://IP:8000/shards

مراحل و خروجی‌ها در --data می‌نشینند؛ اجرای دوباره مراحل تمام‌شده را رد می‌کند.
"""
from __future__ import annotations

import argparse
import collections
import concurrent.futures as cf
import csv
import gzip
import hashlib
import io
import json
import os
import re
import shutil
import sys
import tarfile
import time
import unicodedata
from typing import List, Optional

import numpy as np
import requests
import torch
import torch.nn as nn
from PIL import Image

csv.field_size_limit(1 << 30)

ap = argparse.ArgumentParser()
ap.add_argument("--data", default="/data/zambil", help="پوشه کار: ورودی، چک‌پوینت، خروجی")
ap.add_argument("--shards", required=True, help="مثلا http://38.54.13.36:8000/shards")
ap.add_argument("--per-shard", type=int, default=2000)
ap.add_argument("--decode-workers", type=int, default=0, help="۰ یعنی از روی تعداد هسته‌ها")
ap.add_argument("--text-batch", type=int, default=0, help="۰ یعنی از روی VRAM")
ap.add_argument("--image-batch", type=int, default=0)
ap.add_argument("--epochs", type=int, default=8)
ap.add_argument("--skip-ab", action="store_true", help="آموزش مقایسه‌ای ۵۰۰هزارتایی را نزن")
A = ap.parse_args()

D = A.data.rstrip("/")
CK = D + "/ckpt"
OUT = D + "/out"
for p in (D, CK, OUT):
    os.makedirs(p, exist_ok=True)

DEV = "cuda" if torch.cuda.is_available() else "cpu"
VRAM = torch.cuda.get_device_properties(0).total_memory / 1e9 if DEV == "cuda" else 0.0
NCPU = os.cpu_count() or 4
DEC_WORKERS = A.decode_workers or max(4, min(16, NCPU))
TXT_BS = A.text_batch or (512 if VRAM >= 20 else 256 if VRAM >= 10 else 128)
IMG_BS = A.image_batch or (256 if VRAM >= 20 else 128 if VRAM >= 10 else 64)

ENC_NAME = "BAAI/bge-m3"
SIG_NAME = "google/siglip-base-patch16-256-multilingual"
DIM, HID = 1024, 1024
TITLE_LEN = 64


def log(*a):
    print(time.strftime("[%H:%M:%S]"), *a, flush=True)


log("device=%s vram=%.1fGB cpu=%d decode_workers=%d text_bs=%d image_bs=%d"
    % (DEV, VRAM, NCPU, DEC_WORKERS, TXT_BS, IMG_BS))
if DEV == "cpu":
    log("!! GPU پیدا نشد — این روی CPU روزها طول می‌کشد. متوقف می‌شوم.")
    sys.exit(2)


# ---------------------------------------------------------------- نرمال‌سازی
_AR = {"ي": "ی", "ك": "ک", "ة": "ه"}
_DIGITS = {ord(c): str(i % 10) for i, c in enumerate("۰۱۲۳۴۵۶۷۸۹٠١٢٣٤٥٦٧٨٩")}
_EMOJI = re.compile("[\U0001F000-\U0001FAFF←-⯿️☀-➿]+")
_HARAKAT = re.compile(r"[ً-ْـ]")
_WS = re.compile(r"\s+")


def norm_fa(s: Optional[str]) -> str:
    s = unicodedata.normalize("NFKC", s or "")
    for e in ("n", "r", "t"):
        s = s.replace(chr(92) + e, " ")
    for a, b in _AR.items():
        s = s.replace(a, b)
    s = _EMOJI.sub(" ", s.translate(_DIGITS))
    s = s.replace("‌", " ").replace("‏", "").replace("‎", "")
    return _WS.sub(" ", _HARAKAT.sub("", s)).strip()


def read_csv(p):
    op = gzip.open if p.endswith(".gz") else open
    with op(p, "rt", encoding="utf-8", newline="") as f:
        return list(csv.DictReader(f))


# ---------------------------------------------------------------- داده
CLASSES_PATH = json.load(open(D + "/bsl_classes.json", encoding="utf-8"))
CLASSES = sorted(CLASSES_PATH, key=int)
C2I = {c: i for i, c in enumerate(CLASSES)}
NC = len(CLASSES)
TRAIN = read_csv(D + "/bsl_train.csv")
HOLD = read_csv(D + "/bsl_holdout.csv")
EVAL = read_csv(D + "/zambil_testset200.csv")
IMGS = read_csv(D + "/bsl_img_labels.csv")
IMG_ROW = {r["img_id"]: i for i, r in enumerate(IMGS)}
img_y_all = np.array([C2I[r["category_id"]] for r in IMGS], np.int64)
SHARD_COUNT = (len(IMGS) + A.per_shard - 1) // A.per_shard
log("classes=%d train=%d hold=%d eval=%d images=%d shards=%d"
    % (NC, len(TRAIN), len(HOLD), len(EVAL), len(IMGS), SHARD_COUNT))

EVAL_DIR = D + "/eval_images"
if not os.path.isdir(EVAL_DIR):
    import zipfile
    zipfile.ZipFile(D + "/zambil_testset200_images.zip").extractall(EVAL_DIR)


# ---------------------------------------------------------------- انکودرها
from transformers import AutoModel, AutoProcessor, AutoTokenizer  # noqa: E402

log("loading encoders ...")
tok = AutoTokenizer.from_pretrained(ENC_NAME)
enc = AutoModel.from_pretrained(ENC_NAME, torch_dtype=torch.float16).to(DEV).eval()
sproc = AutoProcessor.from_pretrained(SIG_NAME)
smodel = AutoModel.from_pretrained(SIG_NAME, torch_dtype=torch.float16).to(DEV).eval()
IDIM = smodel.config.vision_config.hidden_size
log("encoders ready, image dim=%d" % IDIM)


@torch.no_grad()
def embed_texts(texts: List[str], tag: str) -> np.ndarray:
    """CLS-pooled و نرمال‌شده، با memmap و ازسرگیری."""
    n = len(texts)
    sig = hashlib.sha1(("|".join(texts[:200]) + str(n)).encode()).hexdigest()[:16]
    npy, mark = "%s/txt_%s_%s.npy" % (CK, tag, sig), "%s/txt_%s_%s.json" % (CK, tag, sig)
    done = 0
    if os.path.exists(mark) and os.path.exists(npy):
        done = int(json.load(open(mark))["done"])
        out = np.lib.format.open_memmap(npy, mode="r+")
        if done >= n:
            log("  %s: از چک‌پوینت کامل" % tag)
            return out
        log("  %s: ادامه از %d/%d" % (tag, done, n))
    else:
        out = np.lib.format.open_memmap(npy, mode="w+", dtype=np.float16, shape=(n, DIM))
    t0 = time.time()
    for i in range(done, n, TXT_BS):
        b = texts[i:i + TXT_BS]
        t = tok(b, padding=True, truncation=True, max_length=TITLE_LEN, return_tensors="pt").to(DEV)
        h = enc(**t).last_hidden_state[:, 0]
        out[i:i + len(b)] = torch.nn.functional.normalize(h.float(), dim=-1).cpu().numpy().astype(np.float16)
        if (i // TXT_BS) % 200 == 0 and i > done:
            r = (i - done) / max(time.time() - t0, 1e-6)
            out.flush()
            json.dump({"done": i}, open(mark, "w"))
            log("  %s %d/%d  %.0f/s  eta %.0f دقیقه" % (tag, i, n, r, (n - i) / max(r, 1) / 60))
    out.flush()
    json.dump({"done": n}, open(mark, "w"))
    log("  %s تمام (%.1f دقیقه)" % (tag, (time.time() - t0) / 60))
    return out


# ---------------------------------------------------------------- تصاویر
IMG_NPY, IMG_MASK, IMG_STATE = CK + "/img_emb.npy", CK + "/img_mask.npy", CK + "/img_state.json"


@torch.no_grad()
def embed_pil(ims):
    px = sproc(images=ims, return_tensors="pt").to(DEV)
    px["pixel_values"] = px["pixel_values"].half()
    f = smodel.get_image_features(**px)
    return torch.nn.functional.normalize(f.float(), dim=-1).cpu().numpy().astype(np.float16)


def fetch_shard(i: int) -> Optional[bytes]:
    for _ in range(3):
        try:
            r = requests.get("%s/shard_%05d.tar" % (A.shards.rstrip("/"), i), timeout=300)
            if r.status_code == 200 and len(r.content) > 1000:
                return r.content
        except Exception as e:
            log("  shard %d: %s" % (i, type(e).__name__))
            time.sleep(3)
    return None


def decode_one(payload):
    """رمزگشایی روی نخ کارگر انجام می‌شود، نه روی نخ اصلی که GPU را تغذیه می‌کند."""
    row, raw = payload
    try:
        return row, Image.open(io.BytesIO(raw)).convert("RGB")
    except Exception:
        return row, None


def run_images():
    if os.path.exists(IMG_STATE) and os.path.exists(IMG_NPY):
        st = json.load(open(IMG_STATE))
        Eimg = np.lib.format.open_memmap(IMG_NPY, mode="r+")
        have = np.load(IMG_MASK)
        sdone = set(st["done"])
        log("تصاویر: ادامه از %d شارد" % len(sdone))
    else:
        Eimg = np.lib.format.open_memmap(IMG_NPY, mode="w+", dtype=np.float16,
                                         shape=(len(IMGS), IDIM))
        have = np.zeros(len(IMGS), bool)
        sdone = set()

    todo = [i for i in range(SHARD_COUNT) if i not in sdone]
    log("تصاویر: %d شارد باقی‌مانده، %d نخ رمزگشا" % (len(todo), DEC_WORKERS))
    t0, n_ok, missing = time.time(), 0, []
    pool = cf.ThreadPoolExecutor(DEC_WORKERS)
    try:
        for k, si in enumerate(todo, 1):
            blob = fetch_shard(si)
            if blob is None:
                missing.append(si)
                continue
            try:
                tf = tarfile.open(fileobj=io.BytesIO(blob))
                jobs = []
                for m in tf.getmembers():
                    row = IMG_ROW.get(m.name[:-4])
                    if row is not None:
                        jobs.append((row, tf.extractfile(m).read()))
            except Exception as e:
                log("  shard %d خراب: %s" % (si, type(e).__name__))
                missing.append(si)
                continue

            ims, rows = [], []
            for row, im in pool.map(decode_one, jobs):
                if im is not None:
                    ims.append(im)
                    rows.append(row)
            for j in range(0, len(ims), IMG_BS):
                chunk, rr = ims[j:j + IMG_BS], rows[j:j + IMG_BS]
                Eimg[rr] = embed_pil(chunk)
                have[rr] = True
            n_ok += len(ims)
            sdone.add(si)
            if k % 25 == 0 or k == len(todo):
                Eimg.flush()
                np.save(IMG_MASK, have)
                json.dump({"done": sorted(sdone)}, open(IMG_STATE, "w"))
                r = n_ok / max(time.time() - t0, 1e-6)
                log("  %d/%d شارد  %d تصویر  %.0f/s  eta %.0f دقیقه"
                    % (k, len(todo), n_ok, r, (len(todo) - k) * A.per_shard / max(r, 1) / 60))
    finally:
        pool.shutdown(wait=True)
        Eimg.flush()
        np.save(IMG_MASK, have)
        json.dump({"done": sorted(sdone)}, open(IMG_STATE, "w"))
    if missing:
        log("!! %d شارد نیامد: %s" % (len(missing), missing[:10]))
    return Eimg, have


# ---------------------------------------------------------------- آموزش سرها
def acc(logits, y, k=1):
    return (logits.topk(k, -1).indices == y[:, None]).any(-1).float().mean().item()


class Head(nn.Module):
    def __init__(self, d, h, k):
        super().__init__()
        self.net = nn.Sequential(nn.Linear(d, h), nn.GELU(), nn.Dropout(0.25), nn.Linear(h, k))

    def forward(self, x):
        return self.net(x)


def train_text_head(X, y, Xv, yv, idx, tag, epochs, cw):
    torch.manual_seed(0)
    h = Head(DIM, HID, NC).to(DEV)
    opt = torch.optim.AdamW(h.parameters(), lr=2e-3, weight_decay=1e-4)
    lf = nn.CrossEntropyLoss(label_smoothing=0.05, weight=cw)
    bs = 2048
    sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=2e-3,
                                                total_steps=epochs * (len(idx) // bs + 1))
    for ep in range(epochs):
        h.train()
        perm = idx[torch.randperm(len(idx), device=idx.device)]
        for i in range(0, len(perm), bs):
            j = perm[i:i + bs]
            xb = X[j].to(DEV, non_blocking=True).float()
            xb = torch.nn.functional.normalize(xb + 0.02 * torch.randn_like(xb), dim=-1)
            loss = lf(h(xb), y[j].to(DEV))
            opt.zero_grad(); loss.backward(); opt.step(); sched.step()
        h.eval()
        with torch.no_grad():
            lv = torch.cat([h(Xv[i:i + 16384].float()) for i in range(0, len(yv), 16384)])
        log("  [%s] epoch %d  holdout top1 %.4f top3 %.4f" % (tag, ep + 1, acc(lv, yv), acc(lv, yv, 3)))
    return h


def heads_of(h):
    return (h.net[0].weight.detach().float().cpu().numpy(),
            h.net[0].bias.detach().float().cpu().numpy(),
            h.net[3].weight.detach().float().cpu().numpy(),
            h.net[3].bias.detach().float().cpu().numpy())


# ---------------------------------------------------------------- اجرا
STAGE = OUT + "/stage.json"
stage = json.load(open(STAGE)) if os.path.exists(STAGE) else {}


def mark(k):
    stage[k] = True
    json.dump(stage, open(STAGE, "w"))


log("=== ۱/۵ امبدینگ عنوان ===")
ytr = np.array([C2I[r["category_id"]] for r in TRAIN], np.int64)
yho = np.array([C2I[r["category_id"]] for r in HOLD], np.int64)
Etr = embed_texts([r["name"] for r in TRAIN], "train")
Eho = embed_texts([r["name"] for r in HOLD], "hold")

log("=== ۲/۵ امبدینگ تصویر ===")
Eimg, have = run_images()
idx_img = np.flatnonzero(have)
img_y = img_y_all[idx_img]
log("تصاویر امبدشده: %d / %d" % (len(idx_img), len(IMGS)))

log("=== ۳/۵ آموزش سر متنی ===")
cnt = np.bincount(ytr, minlength=NC)
cw = (cnt.sum() / np.maximum(cnt, 1)) ** 0.5
cw = torch.tensor(cw / cw.mean(), dtype=torch.float32, device=DEV)
# بردارها روی خود کارت اگر جا باشد؛ وگرنه CPU و انتقال بچ‌به‌بچ
need = Etr.nbytes / 1e9
on_gpu = VRAM >= need + 6
log("بردار متن %.1f گیگ -> %s" % (need, "GPU" if on_gpu else "CPU"))
Xt = torch.from_numpy(np.ascontiguousarray(Etr))
Xt = Xt.to(DEV) if on_gpu else Xt
yt = torch.from_numpy(ytr).to(DEV) if on_gpu else torch.from_numpy(ytr)
Xv = torch.from_numpy(np.ascontiguousarray(Eho)).to(DEV)
yv = torch.from_numpy(yho).to(DEV)
all_idx = torch.arange(len(ytr), device=Xt.device)

if not A.skip_ab:
    g = torch.Generator(device="cpu").manual_seed(1)
    sub = all_idx[torch.randperm(len(all_idx), generator=g).to(all_idx.device)[:500000]]
    log("--- A: %d ردیف ---" % len(sub))
    h_small = train_text_head(Xt, yt, Xv, yv, sub, "500k", A.epochs, cw)
    W1s, b1s, W2s, b2s = heads_of(h_small)
    del h_small
log("--- B: %d ردیف ---" % len(all_idx))
head = train_text_head(Xt, yt, Xv, yv, all_idx, "full", A.epochs, cw)
W1, b1, W2, b2 = heads_of(head)
del Xt, Xv
torch.cuda.empty_cache()
mark("text_head")

log("=== ۴/۵ آموزش سر تصویری ===")
cnt_i = np.bincount(img_y, minlength=NC)
cwi = (cnt_i.sum() / np.maximum(cnt_i, 1)) ** 0.5
cwi = torch.tensor(cwi / cwi.mean(), dtype=torch.float32, device=DEV)
Xi = torch.from_numpy(np.ascontiguousarray(Eimg[idx_img]))
if VRAM >= Xi.nbytes / 1e9 + 6:
    Xi = Xi.to(DEV)
yi = torch.from_numpy(img_y).to(Xi.device)
pv = torch.randperm(len(yi), generator=torch.Generator(device="cpu").manual_seed(0)).to(Xi.device)
nval = max(5000, len(yi) // 20)
vi, ti = pv[:nval], pv[nval:]
Xv_i, yv_i = Xi[vi].to(DEV).float(), yi[vi].to(DEV)
lossf_i = nn.CrossEntropyLoss(label_smoothing=0.05, weight=cwi)


def train_ihead(idx, epochs=10):
    torch.manual_seed(0)
    h = nn.Linear(IDIM, NC).to(DEV)
    o = torch.optim.AdamW(h.parameters(), lr=3e-3, weight_decay=1e-3)
    sc = torch.optim.lr_scheduler.OneCycleLR(o, max_lr=3e-3, total_steps=epochs * (len(idx) // 4096 + 1))
    for _ in range(epochs):
        h.train()
        perm = idx[torch.randperm(len(idx), device=idx.device)]
        for i in range(0, len(perm), 4096):
            j = perm[i:i + 4096]
            loss = lossf_i(h(Xi[j].to(DEV).float()), yi[j].to(DEV))
            o.zero_grad(); loss.backward(); o.step(); sc.step()
    h.eval()
    with torch.no_grad():
        lv = h(Xv_i)
    return h, acc(lv, yv_i), acc(lv, yv_i, 3)


log("--- منحنی یادگیری سر تصویری ---")
curve = []
for frac in (0.1, 0.25, 0.5, 1.0):
    k = max(5000, int(len(ti) * frac))
    _, a1, a3 = train_ihead(ti[:k])
    curve.append((k, a1))
    log("  %8d نمونه (~%4d/کلاس)  top1 %.4f top3 %.4f" % (k, k // NC, a1, a3))
gain = (curve[-1][1] - curve[-2][1]) * 100
log("  رشد بین دو نقطه آخر: %+.2f واحد => %s" % (gain, "هنوز جا دارد" if gain > 0.5 else "اشباع"))
ihead, ia1, ia3 = train_ihead(ti)
IW = ihead.weight.detach().float().cpu().numpy()
Ib = ihead.bias.detach().float().cpu().numpy()
log("سر تصویری: top1 %.4f top3 %.4f" % (ia1, ia3))
mark("image_head")

log("=== ۵/۵ کالیبراسیون، ارزیابی، ذخیره ===")


def _peak(z):
    z = z - z.max(-1, keepdims=True)
    p = np.exp(z); p /= p.sum(-1, keepdims=True)
    return p / np.maximum(p.max(-1, keepdims=True), 1e-9)


def text_logits(e, w=None):
    W1_, b1_, W2_, b2_ = w or (W1, b1, W2, b2)
    h = np.asarray(e, np.float32) @ W1_.T + b1_
    h = h * 0.5 * (1 + np.tanh(0.7978845608 * (h + 0.044715 * h ** 3)))
    return h @ W2_.T + b2_


def fuse(t=None, i=None, wt=1.0, wi=0.35, w=None):
    parts, den = [], 0.0
    if t is not None:
        parts.append(wt * _peak(text_logits(np.atleast_2d(t), w))); den += wt
    if i is not None:
        parts.append(wi * _peak(np.atleast_2d(np.asarray(i, np.float32)) @ IW.T + Ib)); den += wi
    return sum(parts) / den


ev = [r for r in EVAL if r["category_id"] in C2I
      and os.path.exists(EVAL_DIR + "/" + r["product_id"] + ".jpg")]
ey = np.array([C2I[r["category_id"]] for r in ev])
Et = np.asarray(embed_texts([norm_fa(r["name"]) for r in ev], "eval"))
Ei = np.zeros((len(ev), IDIM), np.float16)
for j in range(0, len(ev), IMG_BS):
    ims = [Image.open(EVAL_DIR + "/" + r["product_id"] + ".jpg").convert("RGB")
           for r in ev[j:j + IMG_BS]]
    Ei[j:j + len(ims)] = embed_pil(ims)

best_wi, best = 0.35, -1.0
for wi in (0.0, 0.15, 0.25, 0.35, 0.5, 0.7, 1.0):
    a1 = (fuse(Et, Ei, wi=wi).argmax(-1) == ey).mean()
    log("  wi=%.2f  top1=%.4f%s" % (wi, a1, "  (فقط عنوان)" if wi == 0 else ""))
    if a1 > best:
        best, best_wi = a1, wi
W_IMAGE = best_wi
S = fuse(Et, Ei, wi=W_IMAGE)
pred = S.argmax(-1)
top1 = (pred == ey).mean()
top3 = float(np.mean([ey[i] in np.argsort(-S[i])[:3] for i in range(len(ey))]))
a1_txt = (fuse(Et, wi=0).argmax(-1) == ey).mean()
srt = np.sort(S, axis=-1)
mg = srt[:, -1] - srt[:, -2]
GATE = None
for q in (0.5, 0.6, 0.7, 0.8, 0.9):
    th = float(np.quantile(mg, q))
    m = mg >= th
    a = (pred[m] == ey[m]).mean()
    log("  margin>=%.3f  پوشش=%.1f%%  دقت=%.1f%%" % (th, m.mean() * 100, a * 100))
    if GATE is None and a >= 0.90:
        GATE = th
GATE = GATE if GATE is not None else float(np.quantile(mg, 0.8))

log("")
log("=== نتیجه روی %d محصول واقعی زنبیل ===" % len(ey))
log("  فقط عنوان      : top1 %.1f%%" % (a1_txt * 100))
log("  عنوان + تصویر  : top1 %.1f%%  top3 %.1f%%" % (top1 * 100, top3 * 100))
log("  وزن تصویر      : %.2f   آستانه: %.3f" % (W_IMAGE, GATE))
log("  (مدل قبلی: top1 68.0%  top3 84.5%)")

META = {"inputs": ["title", "image"], "text_encoder": ENC_NAME, "text_dim": DIM,
        "hidden": HID, "image_encoder": SIG_NAME, "image_dim": IDIM, "n_classes": NC,
        "title_max_len": TITLE_LEN, "w_title": 1.0, "w_image": float(W_IMAGE),
        "gate_margin": float(GATE), "text_pooling": "cls", "text_l2_normalized": True,
        "train_rows": int(len(ytr)), "image_rows": int(len(img_y)),
        "eval_top1": float(top1), "eval_top3": float(top3),
        "eval_top1_title_only": float(a1_txt)}
with open(OUT + "/zambil_category_bsl.npz", "wb") as f:
    np.savez_compressed(f, W1=W1, b1=b1, W2=W2, b2=b2, IW=IW, Ib=Ib,
                        classes=np.array(CLASSES),
                        paths=np.array([CLASSES_PATH[c] for c in CLASSES]),
                        meta=np.array([json.dumps(META, ensure_ascii=False)]))
log("ذخیره شد: %s/zambil_category_bsl.npz" % OUT)

# نمایه عنوان: نمونه‌ترین‌های هر دسته، نه تصادفی
TITLES = [r["name"] for r in read_csv(D + "/bsl_img_titles.csv.gz")]
sel = []
for c in range(NC):
    rows = np.flatnonzero(img_y == c)
    if not len(rows):
        continue
    V = np.asarray(Eimg[idx_img[rows]], np.float32)
    cen = V.mean(0)
    n = np.linalg.norm(cen)
    if n:
        cen /= n
    sel.extend(rows[np.argsort(-(V @ cen))[:150]].tolist())
sel = np.array(sel, np.int64)
with open(OUT + "/zambil_title_index.npz", "wb") as f:
    np.savez_compressed(f, vecs=np.asarray(Eimg[idx_img[sel]], np.float16),
                        cats=img_y[sel].astype(np.int16),
                        titles=np.array([TITLES[idx_img[r]] for r in sel]),
                        meta=np.array([json.dumps({"image_encoder": SIG_NAME,
                                                   "image_dim": IDIM, "per_category": 150,
                                                   "n": int(len(sel))}, ensure_ascii=False)]))
log("ذخیره شد: %s/zambil_title_index.npz  (%d بردار)" % (OUT, len(sel)))
json.dump({"top1": float(top1), "top3": float(top3), "top1_title_only": float(a1_txt),
           "w_image": float(W_IMAGE), "gate": float(GATE),
           "finished_at": time.strftime("%Y-%m-%d %H:%M:%S")},
          open(OUT + "/result.json", "w"), ensure_ascii=False, indent=1)
mark("done")
log("=== تمام ===")
