diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..f162c0f --- /dev/null +++ b/.gitignore @@ -0,0 +1,5 @@ +__pycache__/ +*.py[cod] +models/ +output/ +.pytest_cache/ diff --git a/README.md b/README.md index 9c58803..42a3344 100644 --- a/README.md +++ b/README.md @@ -1,10 +1,10 @@ -# comfyui-lsnet +# comfyui-kaloscope ### *「我这双眼睛,能将黑暗看得一清二楚」* —— 宇智波佐助 -![阿吗特拉斯](https://github.com/user-attachments/assets/a9d16b72-b577-4458-bc10-604eb82fefea) +![阿妈特拉斯](https://github.com/user-attachments/assets/a9d16b72-b577-4458-bc10-604eb82fefea) > “*Kaloscope*”(万花筒)致敬万花筒写轮眼,象征忍术(画风)复刻能力 @@ -12,21 +12,31 @@ ## 核心能力 -基于 *LSNet* 技术核心,本工具聚焦两大核心场景: +支持 *LSNet* 和 *DINOv3* 模型,本工具聚焦以下场景: -1. **画风分类**:识别单幅作品的风格属性,完成风格相似的标签匹配; +1. **画风分类**:识别单幅作品的风格属性,完成风格相似的标签匹配 2. **画风聚类**:自动对多组作品按风格特征进行归类聚合,筛选出风格相似的作品群体,实现批量风格整理与分析,提升处理效率 +3. **特征提取**:输出整张图片、patch tokens、空间特征图或中间层特征,方便连接其他分析节点 + +4. **关系制图**:把多张图片的相似关系画成关系网络、距离热图或聚类散点,也可以查看特征统计、近邻排行和 patch 能量分布 + ## 第一步:下载必要文件 -前往 Hugging Face 仓库,下载两个核心文件: +前往 Hugging Face 或 ModelScope 仓库,下载模型对应的文件: -* `best_checkpoint.pth`(*LSNet 模型权重文件*,决定画风分类与聚类的精度) +* `best_checkpoint.pth` / `best.pt` / `model.safetensors`(*模型权重文件*,支持 `.pt`、`.pth`、`.ckpt`、`.safetensors`) -* `class_mapping.csv`(*风格类别映射配置文件*,适配不同画风标签的识别与归类) +* `class_mapping.csv`(*风格类别映射配置文件*,可选,用于把类别编号转换成画师或风格名称) -* `config.json`(*LSNet 模型配置文件*,如果有的话请下载) +* `config.json`(用于识别模型架构,如果有的话请下载) + +### v3版本 + +huggingface仓库地址:暂未开源 + +或者在modelscope下载:暂未开源 ### v2版本 @@ -44,9 +54,9 @@ huggingface仓库地址:https://huggingface.co/heathcliff01/Kaloscope/tree/mai ### 1. 创建目录结构 -在 ComfyUI 的`models`目录下,新建名为`lsnet`的文件夹(专门存放 LSNet 相关模型文件); +在 ComfyUI 的`models`目录下,新建名为`kaloscope`的文件夹(存放模型文件); -进入`lsnet`文件夹后,可随意创建一个子文件夹(如 “checkpoints”“kaloscope” 等,名称无强制要求,用于归类核心文件)。 +进入`kaloscope`文件夹后,可随意创建一个子文件夹(如 “checkpoints”“kaloscope” 等,名称无强制要求,用于归类核心文件)。 目录结构示例: @@ -56,67 +66,367 @@ ComfyUI/ └── models/ - └── lsnet/ + └── kaloscope/ - └── 子文件夹名称/ # 例:"sharingan" 或 “kaloscope” + └── 子文件夹名称/ # 例:"kaloscope-v2” - ├── best_checkpoint.pth + ├── best_checkpoint.pth # 或 best.pt / model.safetensors - └── class_mapping.csv + ├── class_mapping.csv # 可选 - └── config.json # 如果有 + └── config.json # 填写模型架构 ``` -相关操作截图: - -* 新建`lsnet`文件夹: - -![新建lsnet文件夹](https://github.com/user-attachments/assets/d959be3c-156c-4c54-9076-f9f5a25000a9) - -* 在`lsnet`内创建子文件夹: - -![创建子文件夹](https://github.com/user-attachments/assets/f64d8e9c-8047-424b-b9b0-8a6ec1732ef0) - ### 2. 安装依赖 -将上述两个(或三个)文件放入子文件夹后,执行以下命令安装 LSNet 运行所需依赖(webui插件可以跳过这一步,会自动安装依赖): +将模型权重、配置和可选类别映射放入子文件夹后,在插件目录使用 ComfyUI 的 Python 环境安装依赖(webui插件可以跳过这一步,会自动安装依赖): ``` -pip install -r requirements.txt --upgrade +python -m pip install -r requirements.txt ``` ## 第三步:启动 ComfyUI 并使用 1. 按常规方式启动 ComfyUI -2. 在画风分析工作流中,调用 LSNet 相关节点,即可触发 “画风分类” 或 “画风聚类” 功能 +2. 在画风分析工作流中,调用 **Kaloscope** 分类下的节点,即可使用画风分类、特征提取、画风聚类和关系制图 -> ps:目前版本该插件已经可以单独启动和作为webui插件启动,单独启动在项目根目录运行python -m scripts.app,模型路径为根目录下models/lsnet文件夹内,webui插件启动和comfyui差不多 ### 使用示例 -![基础推理界面(含LSNet核心节点)](https://github.com/user-attachments/assets/28cc2820-ff5d-4290-8ac2-339763947e91) +### 1. 分类与画风比较 -聚类暂时不作示例,内有其他节点供开发者使用 +* **Kaloscope Model Loader**:选择模型文件夹,输出模型 + +* **Kaloscope Artist Inference**:接图片和模型,输出标签与 JSON。无分类头时标签为空,JSON 输出特征 + +* **Kaloscope Artist Similarity**:比较一张查询图片与多张参考图片,输出余弦相似度 + +* **Kaloscope Common Features**:接一组参考图片,输出该组平均特征 `[D]` + +* **Kaloscope Feature Comparison**:接查询图片与最多三组平均特征,输出各组相似度和最接近的组 + +* **Kaloscope Clustering**:接最多三组 `[B,D]` 特征,支持 KMeans、DBSCAN、hierarchical 聚类,可输出 PCA/t-SNE 图 + +* **Kaloscope Image Connector**:将三路同尺寸图片组成一个批次,每路若为批次则取第一张。更多图片可以使用 ComfyUI 的图片批次组合节点 + +分组比较时,每组图片分别接 **Common Features**,再把平均特征接到 **Feature Comparison** 的 `group_1` / `group_2` / `group_3` + +聚类时把 **Extract Features** 的输出接到 **Clustering**;每张图片都需要保留自己的特征,不能用组均值替代。KMeans/hierarchical 的 `n_clusters` 要按图片数量设置,DBSCAN 使用 `eps` / `min_samples` 控制 + +### 2. 提取特征给下游使用 + +**Kaloscope Extract Features** 接图片批次和模型,最后输出一个 CPU float32 TENSOR。在 `output_type` 中选择需要的特征 + +| output_type | 输出形状 | 用途 | +| --- | --- | --- | +| `default` | `[B,F]` | 按模型配置选择骨干或投影特征,适合先做画风比较 | +| `backbone` | `[B,D]` 或 `[B,2D]` | 按模型池化配置输出骨干特征 | +| `cls` | `[B,D]` | CLS token,全局特征 | +| `mean` | `[B,D]` | patch tokens 的均值 | +| `cls_mean` | `[B,2D]` | CLS 与 patch 均值拼接 | +| `projector` | `[B,P]` | 模型投影层输出 | +| `patch_tokens` | `[B,N,D]` | 局部 patch 特征 | +| `patch_map` | `[B,D,H,W]` | 空间特征图 | +| `storage_tokens` | `[B,R,D]` | storage/register tokens | +| `all_tokens` | `[B,1+R+N,D]` | CLS、storage、patch tokens 按顺序拼接 | +| `prenorm` | `[B,1+R+N,D]` | 最后一层 LayerNorm 之前的完整 tokens | +| `intermediate_cls` | `[B,L,D]` | 中间层 CLS | +| `intermediate_mean` | `[B,L,D]` | 中间层 patch 均值 | +| `intermediate_cls_mean` | `[B,L,2D]` | 中间层 CLS 与 patch 均值拼接 | +| `intermediate_patch_tokens` | `[B,L,N,D]` | 中间层 patch tokens | +| `intermediate_patch_map` | `[B,L,D,H,W]` | 中间层空间特征图 | +| `intermediate_storage_tokens` | `[B,L,R,D]` | 中间层 storage tokens | +| `intermediate_all_tokens` | `[B,L,1+R+N,D]` | 中间层完整 tokens | +| `intermediate_prenorm` | `[B,L,1+R+N,D]` | 中间层未归一化 tokens | + +`B` 是图片数,`D` 是通道数,`N` 是 patch 数,`R` 是 storage token 数,`L` 是选取的层数,`P` 是投影维度,`F` 是模型默认特征维度。表中的 token 形状以 ViT 为例,实际维度随架构和输入尺寸变化 + +* `layers`:填写中间层编号,例如 `-1` 为最后一层,`8,9,10,11` 或 `-4,-3,-2,-1` 为 ViT-B 的最后四层。层编号从 0 开始,输出顺序与填写顺序一致,单层也保留 `L=1` + +* `intermediate_norm`:是否对中间层应用模型的 LayerNorm,默认开启;`intermediate_prenorm` 始终不应用 LayerNorm + +* LSNet 支持 `default` / `backbone`;没有 projector 或 storage tokens 的模型不能选择对应输出 + +* ConvNeXt 的 CLS 表示全局池化,`prenorm` 为未归一化的空间 tokens。不同 stage 的形状可能不同,请一次选择一个 stage + +> ps:分类头的池化方式不受这里的选择影响。同一轮图片比较请使用同一个模型、特征类型和预处理设置。 + +### 3. 从特征生成关系图和分析图 + +想看几张图片之间的关系,可以按这个方式连接: + +```text +图片批次 + Kaloscope Model Loader + ↓ + Kaloscope Extract Features + ↓ TENSOR + Kaloscope Feature Analysis + ↓ IMAGE + Preview Image / Save Image +``` + +**Kaloscope Feature Analysis** 直接使用已有特征,不再执行推理,输出 + +* `visualization`:图像,可以连接预览或保存节点 + +* `analysis_json`:距离、相似度、近邻、簇标签、投影坐标和统计结果,方便下游读取 + +* `distance_matrix`:`[B,B]` 距离 TENSOR + +节点在 **Kaloscope/Analysis** 下。同一份特征可以接多个分析节点,分别生成不同图表。可选 `images` 只用来显示缩略图,图片数量和顺序要与特征一致 + +也可以用 **Kaloscope Image Analysis** 直接接图片与模型;它提取特征后制图,同时输出 `features`,可以再连接其他分析节点 + +> ps:ComfyUI 中组成图片批次前需要统一尺寸。想看类似 KMeans 的关系分组,先选择 `relationship_graph`,再设置 `cluster_method=kmeans` 和 `n_clusters` + +| chart_type | 图表用途 | +| --- | --- | +| `relationship_graph` | 近邻关系网络,MDS 布局、聚类颜色/标记,连线数字表示原始特征距离 | +| `distance_heatmap` | 两两距离热图 | +| `similarity_heatmap` | 两两余弦相似度热图 | +| `pca_scatter` | PCA 二维散点,显示解释方差比例 | +| `mds_scatter` | 近似保持距离的二维散点 | +| `tsne_scatter` | t-SNE 邻域结构散点 | +| `dendrogram` | 层次聚类关系树 | +| `nearest_neighbors` | 指定图片的最近邻排行 | +| `distance_distribution` | 图片对的距离分布与簇内/簇间距离 | +| `silhouette` | 各图片轮廓系数及平均值 | +| `cluster_sizes` | 各簇数量,包含 DBSCAN 噪声 | +| `pca_variance` | PCA 方差解释率、累计比例与有效秩 | +| `feature_statistics` | 特征范数、平均绝对值和标准差 | +| `feature_heatmap` | 图片与高方差特征维度的数值热图 | +| `dimension_correlation` | 特征维度之间的 Pearson 相关 | +| `cluster_centroid_heatmap` | 各簇平均特征热图 | +| `outlier_scores` | 最近邻距离均值,用来查看孤立程度 | +| `patch_energy` | patch 特征 L2 范数的空间分布 | + +常用参数: + +* `metric`:`cosine`、`euclidean`、`manhattan`。余弦距离为 `1-cosine_similarity`,相似度热图始终显示余弦相似度; + +* `normalize`:是否在距离和聚类前按图片做 L2 归一化,默认开启。原始特征统计与 patch 能量使用归一化之前的输入; + +* `cluster_method`:`kmeans`、`agglomerative`、`dbscan`、`none`。KMeans 使用欧氏向量目标,agglomerative/DBSCAN 使用所选距离; + +* `n_clusters`:KMeans/agglomerative 的簇数;图片或不同向量数量不足时自动减少。`dbscan_eps` / `dbscan_min_samples` 为 DBSCAN 参数; + +* `top_k`:关系连线、近邻排行和孤立得分的邻居数;`reference_index` 选择查询图片,从 0 开始; + +* `labels`:每行一个名字或 JSON 数组,顺序与输入图片一致; + +* `seed` / `perplexity`:随机种子与 t-SNE 参数; + +* `max_dimensions`:特征热图最多显示多少个高方差维度;`heatmap_order` 选择 `cluster` 或 `input` 排序; + +* `width` / `height`:输出分辨率,范围 512–4096 像素;`grid_width` 指定 patch 网格列数,0 自动推断正方形网格 + +局部和多层特征需要按输出形状选择 `tensor_layout`: + +| 特征类型 | tensor_layout | +| --- | --- | +| 全局向量 `[B,D]` | `vectors` 或 `auto` | +| patch/storage/all tokens `[B,N,D]` | `tokens` | +| patch_map `[B,D,H,W]` | `spatial` | +| 中间层全局向量 `[B,L,D]` | `layer_vectors` | +| 中间层 tokens `[B,L,N,D]` | `layer_tokens` | +| 中间层 patch_map `[B,L,D,H,W]` | `layer_spatial` | + +`layer_index` 选择输入特征中的层位置,默认最后一层;`layer_pooling=mean` 对选取的层取均值。`token_pooling` 可选 `mean` 或 `flatten`,`tensor_layout=flatten` 可将每张图片的其余维度直接展平 + +Image Analysis 和界面缓存会根据特征类型处理布局。直接传 TENSOR 时,四维输入请明确选择 `spatial` 或 `layer_tokens` + +## 第四步:作为 WebUI 插件或单独启动 + +### 1. WebUI 插件 + +将插件放到 WebUI 的 `extensions/comfyui-kaloscope/`,重启后打开 **Kaloscope** 页签。支持 AUTOMATIC1111 及兼容其扩展接口的 WebUI + +模型放在 WebUI 的 `models/kaloscope/<子文件夹>/`,权重、配置和类别映射的放法与 ComfyUI 相同 + +### 2. 单独启动 + +在项目根目录安装界面依赖,然后启动: + +```bash +python -m pip install -r requirements-webui.txt +python -m scripts.app +``` + +浏览器打开 `http://127.0.0.1:7860`。也可以运行 `python scripts/app.py` 或双击 `单独启动.bat`。 + +模型默认放在项目根目录的 `models/kaloscope/<子文件夹>/`。需要更改模型目录或端口时: + +```bash +python -m scripts.app --models-dir D:/models --host 127.0.0.1 --port 7860 +``` + +这里的 `D:/models` 是包含 `kaloscope/` 的根目录,也可通过环境变量 `KALOSCOPE_MODELS_DIR` 指定。 + +### 3. 界面里怎么用 + +WebUI 和独立启动使用同一套界面: + +1. **Inference**:上传单张图片,选择模型、设备、Top K 和阈值,点击 Infer。有分类头输出分类,无分类头输出特征 + +2. **Features & Analysis**:上传多张图片,选择特征类型、中间层和批次大小,点击「提取并缓存特征」 + +3. 选择图表和分析参数,点击「从缓存生成图表」,可以反复换图表,不需要重新提取 + +4. 下载 PNG、分析 JSON、距离矩阵 CSV,也可以下载 `features.npz`,下次直接导入缓存 + +5. 「共同特征 / 相似度 / 分组比较」使用同一份缓存,选择查询图片编号;分组比较时,每张图片填写一行组名 + +`mode=auto` 自动选择分类或特征;`cluster` 只提取特征;`classify` / `both` 分别用于分类、分类加特征,需要模型有分类头 + +## 第五步:命令行和 API 用法 + +### 1. 命令行推理 + +在项目根目录运行,例如模型放在 `models/kaloscope/sharingan/`: + +```bash +python inference_artist.py --checkpoint models/kaloscope/sharingan/best.pt --input example.png --device cuda --mode auto --output output +``` + +`--input` 也可以填写图片目录。`--mode cluster` 提取特征,`--output-type` 选择特征类型,`--layers` 选择中间层,`--no-intermediate-norm` 关闭中间层 LayerNorm。 + +分类结果保存为 JSON;提取特征时同时保存 `features.npz`,批量提取还会保存 `features.npy` 与图片名称列表 + +### 2. 命令行制图 + +```bash +# 从一组图片提取一次特征,再生成关系图 +python analysis_cli.py --input images --model-dir models/kaloscope/sharingan --output outputs --device cuda + +# 用缓存改画距离热图,不加载模型 +python analysis_cli.py --features outputs/features.npz --chart-type distance_heatmap --output outputs + +# 使用 patch 特征,一次提取后生成全部 18 类图 +python analysis_cli.py --input images --model-dir models/kaloscope/sharingan --output-type patch_tokens --all-charts --output outputs +``` + +每种图会保存 PNG、JSON 和距离矩阵 CSV,参数与界面对应,例如 `--metric euclidean --no-normalize --cluster-method dbscan --dbscan-eps 0.5`。 + +还可以从缓存获取共同特征、相似度或分组比较: + +```bash +python analysis_cli.py --features outputs/features.npz --operation common_features --output outputs +python analysis_cli.py --features outputs/features.npz --operation similarity --reference-index 0 --output outputs +python analysis_cli.py --features outputs/features.npz --operation compare_groups --groups groups.txt --output outputs +``` + +`groups.txt` 每行一个组名,与图片顺序一致。`--options-json options.json` 可以读取分析参数,命令行显式填写的参数优先;所有参数可通过 `python analysis_cli.py --help` 查看 + +> ps:界面、API 和命令行共用 `.npz` 缓存;`.npy` 只包含数组。生成全部图表需要 patch tokens/patch_map,因为全局向量不能生成 patch 能量图 + +### 3. API + +WebUI 和独立服务都提供 `/kaloscope/v1/` API。独立服务打开 `/docs` 可以查看请求参数 + +| 接口 | 用途 | +| --- | --- | +| `GET /kaloscope/v1/models` | 查看模型目录、特征类型和图表类型 | +| `POST /kaloscope/v1/infer` | 单图分类或提取特征 | +| `POST /kaloscope/v1/features` | 批量提取,返回可复用的特征缓存 | +| `POST /kaloscope/v1/analyze` | 从特征、缓存或图片批次生成图表 | +| `POST /kaloscope/v1/feature-tools` | 共同特征、相似度和分组比较 | + +单图推理: + +```json +{ + "input_image": "<图片的Base64>", + "model_name": "sharingan", + "device": "cuda", + "mode": "auto", + "top_k": 5, + "threshold": 0.0 +} +``` + +`/infer` 返回 `results` 和 `info`。分类列表包含 `class_id`、`class_name`、`probability`;特征保存在 `results.features`。可以填写 `output_type`、`layers`、`intermediate_norm` 选择特征 + +批量提取 `/features`: + +```json +{ + "input_images": ["<图片1的Base64>", "<图片2的Base64>"], + "model_name": "sharingan", + "device": "cuda", + "output_type": "patch_tokens", + "batch_size": 2, + "labels": ["image1", "image2"] +} +``` + +返回 `cache_base64`、`shape`、`labels`、`output_type`。把 `cache_base64` 传给 `/analyze` 就可以制图: + +```json +{ + "cache_base64": "<上一步返回的特征缓存>", + "chart_type": "relationship_graph", + "options": { + "metric": "cosine", + "cluster_method": "kmeans", + "n_clusters": 2, + "top_k": 1 + } +} +``` + +返回 PNG 的 `image_base64`、分析对象 `analysis`、`distance_matrix` 和可复用缓存。`cache_base64` 解码后是完整的 `features.npz` 文件,可以导入界面或用于命令行 + +`/analyze` 也可以接 `features: [[...], [...]]` 数组,填写对应 `output_type` 和 `labels`;或者接 `image_batch`,结构与 `/features` 的请求相同。三种输入选择一种即可,`thumbnail_images` 是可选缩略图,不参与特征推理 + +共同特征、相似度和分组比较可以这样调用 `/feature-tools`: + +```json +{ + "cache_base64": "<特征缓存>", + "operation": "compare_groups", + "reference_index": 0, + "groups": ["artist_a", "artist_a", "artist_b", "artist_b"] +} +``` + +`groups` 数量与缓存图片数一致。`common_features` 返回平均向量与样本数;`similarity` 返回查询图片与批次内所有图片的相似度,包含自身;`compare_groups` 返回组名、相似度和最接近的组 + +> ps:图片 Base64 不带 `data:image/...;base64,` 前缀。`options` 使用分析节点的同名参数,API 的 `labels` 填字符串数组;共同特征和分组比较也支持 `tensor_layout`、`layer_index`、`layer_pooling`、`token_pooling` ### 致谢 感谢 [@heathcliff01](https://huggingface.co/heathcliff01) 训练模型 -### 训练代码 +### lsnet训练代码 https://github.com/spawner1145/lsnet-test.git +### dinov3训练代码 + +https://github.com/Chenkin-x/kaloscope-dinov3.git + ## Citation ```BibTeX @misc{wang2025lsnetlargefocussmall, - title={LSNet: See Large, Focus Small}, - author={Ao Wang and Hui Chen and Zijia Lin and Jungong Han and Guiguang Ding}, - year={2025}, - eprint={2503.23135}, - archivePrefix={arXiv}, - primaryClass={cs.CV}, - url={https://arxiv.org/abs/2503.23135}, + title={LSNet: See Large, Focus Small}, + author={Ao Wang and Hui Chen and Zijia Lin and Jungong Han and Guiguang Ding}, + year={2025}, + eprint={2503.23135}, + archivePrefix={arXiv}, + primaryClass={cs.CV}, + url={https://arxiv.org/abs/2503.23135}, } + +@misc{simeoni2025dinov3, + title={{DINOv3}}, + author={Sim{\'e}oni, Oriane and Vo, Huy V. and Seitzer, Maximilian and Baldassarre, Federico and Oquab, Maxime and Jose, Cijo and Khalidov, Vasil and Szafraniec, Marc and Yi, Seungeun and Ramamonjisoa, Micha{\"e}l and Massa, Francisco and Haziza, Daniel and Wehrstedt, Luca and Wang, Jianyuan and Darcet, Timoth{\'e}e and Moutakanni, Th{\'e}o and Sentana, Leonel and Roberts, Claire and Vedaldi, Andrea and Tolan, Jamie and Brandt, John and Couprie, Camille and Mairal, Julien and J{\'e}gou, Herv{\'e} and Labatut, Patrick and Bojanowski, Piotr}, + year={2025}, + eprint={2508.10104}, + archivePrefix={arXiv}, + primaryClass={cs.CV}, + url={https://arxiv.org/abs/2508.10104}, +} + ``` diff --git a/__init__.py b/__init__.py index dbff772..65844be 100644 --- a/__init__.py +++ b/__init__.py @@ -2,43 +2,28 @@ import os import sys import json import torch -import torch.nn.functional as F from PIL import Image import numpy as np -from pathlib import Path -from typing import Dict, Optional sys.path.append(os.path.dirname(__file__)) import folder_paths -from timm.data import resolve_data_config -from timm.data.transforms_factory import create_transform -from timm.models import create_model - -from lsnet_model import lsnet_artist # noqa: F401 - -from inference_artist import ( - load_checkpoint_state, - normalize_state_dict_keys, - resolve_num_classes, - resolve_feature_dim, - load_class_mapping -) +from model_loading import FEATURE_OUTPUTS, load_model_bundle, model_folders +from inference_artist import classify_image, extract_features +from feature_analysis import CHART_TYPES, TENSOR_LAYOUTS, analyze_features +from backend_lsnet.analysis import extract_batch, output_layout from sklearn.cluster import KMeans, DBSCAN, AgglomerativeClustering from sklearn.manifold import TSNE from sklearn.decomposition import PCA import matplotlib.pyplot as plt -class LSNetModelLoader: +class KaloscopeModelLoader: @classmethod def INPUT_TYPES(s): - base_dir = os.path.join(folder_paths.models_dir, 'lsnet') - subfolders = [] - if os.path.exists(base_dir): - subfolders = [f for f in os.listdir(base_dir) if os.path.isdir(os.path.join(base_dir, f))] - + subfolders = sorted(model_folders(folder_paths.models_dir)) + return { "required": { "model_folder": (subfolders, {"default": subfolders[0] if subfolders else ""}), @@ -46,77 +31,24 @@ class LSNetModelLoader: } } - RETURN_TYPES = ("LSNET_MODEL",) + RETURN_TYPES = ("KALOSCOPE_MODEL",) RETURN_NAMES = ("model",) FUNCTION = "load" - CATEGORY = "LSNet" + CATEGORY = "Kaloscope" def load(self, model_folder, device): - base_dir = os.path.join(folder_paths.models_dir, 'lsnet') - model_dir = os.path.join(base_dir, model_folder) - checkpoint_path = os.path.join(model_dir, "best_checkpoint.pth") - csv_path = os.path.join(model_dir, "class_mapping.csv") + folders = model_folders(folder_paths.models_dir) + if model_folder not in folders: + raise FileNotFoundError(f"Model folder not found: {model_folder}") + return (load_model_bundle(folders[model_folder], device=device),) - if not os.path.exists(checkpoint_path): - raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}") - if not os.path.exists(csv_path): - raise FileNotFoundError(f"Class mapping CSV not found: {csv_path}") - class_mapping = load_class_mapping(csv_path) - state_dict = load_checkpoint_state(checkpoint_path) - state_dict = normalize_state_dict_keys(state_dict) - num_classes = resolve_num_classes(None, class_mapping, state_dict) - feature_dim = resolve_feature_dim(None, state_dict) - - # 自动从config.json读取model类型 - config_path = os.path.join(model_dir, "config.json") - model_type = 'lsnet_xl_artist' # 默认值 - if os.path.exists(config_path): - try: - with open(config_path, 'r', encoding='utf-8') as f: - config = json.load(f) - if 'model' in config and config['model'] in ['lsnet_t_artist', 'lsnet_s_artist', 'lsnet_b_artist', 'lsnet_l_artist', 'lsnet_xl_artist', 'lsnet_xl_artist_448']: - model_type = config['model'] - print(f"Model type loaded from config: {model_type}") - except Exception as e: - print(f"Warning: Failed to load config.json: {e}") - - model = create_model( - model_type, - pretrained=False, - num_classes=num_classes, - feature_dim=feature_dim, - ) - model.load_state_dict(state_dict, strict=False) - model.to(device) - model.eval() - - # 根据模型配置动态设置输入大小 - from lsnet_model.lsnet_artist import default_cfgs_artist - input_size = 224 # 默认值 - if model_type in default_cfgs_artist: - model_cfg = default_cfgs_artist[model_type] - configured_input_size = model_cfg.get('input_size', (3, 224, 224))[1] # 获取高度(假设正方形) - input_size = configured_input_size - print(f"Auto-setting input_size to {input_size} for model {model_type}") - - config = resolve_data_config({'input_size': (3, input_size, input_size)}, model=model) - transform = create_transform(**config) - model_bundle = { - 'model': model, - 'transform': transform, - 'class_mapping': class_mapping, - 'device': device - } - - return (model_bundle,) - -class LSNetArtistInferenceNode: +class KaloscopeArtistInferenceNode: @classmethod def INPUT_TYPES(s): return { "required": { "image": ("IMAGE",), - "model": ("LSNET_MODEL",), + "model": ("KALOSCOPE_MODEL",), "top_k": ("INT", {"default": 5, "min": 1, "max": 100}), "threshold": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0}), } @@ -125,7 +57,7 @@ class LSNetArtistInferenceNode: RETURN_TYPES = ("STRING", "STRING") RETURN_NAMES = ("tag_string", "json_output") FUNCTION = "process" - CATEGORY = "LSNet" + CATEGORY = "Kaloscope" def process(self, image, model, top_k, threshold): model_bundle = model @@ -142,27 +74,11 @@ class LSNetArtistInferenceNode: # Preprocess image image_tensor = transform(pil_image).unsqueeze(0) # Add batch dimension - # Classify - with torch.no_grad(): - image_tensor = image_tensor.to(device) - logits = model(image_tensor, return_features=False) - probs = F.softmax(logits, dim=-1) - top_probs, top_indices = torch.topk(probs, k=min(top_k, probs.size(-1)), dim=-1) - - results = [] - for prob, idx in zip(top_probs[0].cpu().numpy(), top_indices[0].cpu().numpy()): - if prob >= threshold: - class_id = int(idx) - class_name = class_mapping.get(class_id, f"Class {class_id}") - results.append({ - 'class_id': class_id, - 'class_name': class_name, - 'probability': float(prob) - }) - - # Limit to top_k if more results after filtering - if len(results) > top_k: - results = results[:top_k] + if not model_bundle['has_classifier']: + features = extract_features(model, image_tensor, device)[0].tolist() + return ('', json.dumps({'features': features, 'feature_dim': len(features), + 'feature_source': model_bundle['feature_source']}, ensure_ascii=False)) + results = classify_image(model, image_tensor, device, class_mapping, top_k, threshold) # Prepare outputs tags = [res['class_name'] for res in results] @@ -172,21 +88,21 @@ class LSNetArtistInferenceNode: return (tag_string, json_output) -class LSNetArtistSimilarityNode: +class KaloscopeArtistSimilarityNode: @classmethod def INPUT_TYPES(s): return { "required": { "processed_image": ("IMAGE",), "reference_images": ("IMAGE",), - "model": ("LSNET_MODEL",), + "model": ("KALOSCOPE_MODEL",), } } RETURN_TYPES = ("STRING",) RETURN_NAMES = ("similarity_json",) FUNCTION = "process" - CATEGORY = "LSNet" + CATEGORY = "Kaloscope" def process(self, processed_image, reference_images, model): model_bundle = model @@ -228,20 +144,20 @@ class LSNetArtistSimilarityNode: return (json_output,) -class LSNetCommonFeaturesNode: +class KaloscopeCommonFeaturesNode: @classmethod def INPUT_TYPES(s): return { "required": { "reference_images": ("IMAGE",), - "model": ("LSNET_MODEL",), + "model": ("KALOSCOPE_MODEL",), } } RETURN_TYPES = ("TENSOR",) RETURN_NAMES = ("common_features",) FUNCTION = "process" - CATEGORY = "LSNet" + CATEGORY = "Kaloscope" def process(self, reference_images, model): model_bundle = model @@ -269,10 +185,10 @@ class LSNetCommonFeaturesNode: if references: common_features = np.mean(np.array(references), axis=0) else: - common_features = np.zeros(384) + common_features = np.zeros(model_bundle['feature_dim']) return (torch.tensor(common_features),) -class LSNetClusteringNode: +class KaloscopeClusteringNode: @classmethod def INPUT_TYPES(s): return { @@ -295,7 +211,7 @@ class LSNetClusteringNode: RETURN_TYPES = ("STRING", "IMAGE") RETURN_NAMES = ("clustering_json", "visualization") FUNCTION = "cluster" - CATEGORY = "LSNet" + CATEGORY = "Kaloscope" def cluster(self, method, n_clusters, eps, min_samples, visualize, viz_method, perplexity, group_1=None, group_2=None, group_3=None): groups = [] @@ -356,9 +272,9 @@ class LSNetClusteringNode: fig = plt.gcf() fig.canvas.draw() - img_array = np.frombuffer(fig.canvas.tostring_rgb(), dtype=np.uint8) - img_array = img_array.reshape(fig.canvas.get_width_height()[::-1] + (3,)) - pil_image = Image.fromarray(img_array) + img_array = np.frombuffer(fig.canvas.buffer_rgba(), dtype=np.uint8) + img_array = img_array.reshape(fig.canvas.get_width_height()[::-1] + (4,)) + pil_image = Image.fromarray(img_array[:, :, :3]) plt.close() viz_tensor = torch.from_numpy(np.array(pil_image)).float() / 255.0 @@ -369,13 +285,13 @@ class LSNetClusteringNode: return (json_output, viz_tensor) -class LSNetFeatureComparisonNode: +class KaloscopeFeatureComparisonNode: @classmethod def INPUT_TYPES(s): return { "required": { "image": ("IMAGE",), - "model": ("LSNET_MODEL",), + "model": ("KALOSCOPE_MODEL",), }, "optional": { "group_1": ("TENSOR",), @@ -387,7 +303,7 @@ class LSNetFeatureComparisonNode: RETURN_TYPES = ("STRING",) RETURN_NAMES = ("comparison_json",) FUNCTION = "compare" - CATEGORY = "LSNet" + CATEGORY = "Kaloscope" def compare(self, image, model, group_1=None, group_2=None, group_3=None): model_bundle = model @@ -429,7 +345,7 @@ class LSNetFeatureComparisonNode: return (json_output,) -class LSNetArtistImageConnector: +class KaloscopeArtistImageConnector: @classmethod def INPUT_TYPES(s): return { @@ -443,7 +359,7 @@ class LSNetArtistImageConnector: RETURN_TYPES = ("IMAGE",) RETURN_NAMES = ("stacked_images",) FUNCTION = "connect" - CATEGORY = "LSNet" + CATEGORY = "Kaloscope" def connect(self, image_1, image_2, image_3): def normalize_image(img): @@ -458,22 +374,116 @@ class LSNetArtistImageConnector: stacked = torch.cat([img1, img2, img3], dim=0) return (stacked,) +class KaloscopeExtractFeaturesNode: + @classmethod + def INPUT_TYPES(cls): + return { + 'required': {'image': ('IMAGE',), 'model': ('KALOSCOPE_MODEL',)}, + 'optional': { + 'output_type': (list(FEATURE_OUTPUTS), {'default': 'default'}), + 'layers': ('STRING', {'default': '-1', 'tooltip': 'Intermediate layer indices, e.g. -1 or 8,9,10,11; negative indices count from the end.'}), + 'intermediate_norm': ('BOOLEAN', {'default': True, 'tooltip': 'Apply model LayerNorm to intermediate features; prenorm always skips it.'}), + }, + } + + RETURN_TYPES = ('TENSOR',) + RETURN_NAMES = ('features',) + FUNCTION = 'extract' + CATEGORY = 'Kaloscope' + + @torch.inference_mode() + def extract(self, image, model, output_type='default', layers='-1', intermediate_norm=True): + images = image if image.ndim == 4 else image.unsqueeze(0) + pil_images = [Image.fromarray((img * 255).clamp(0, 255).byte().cpu().numpy()) for img in images] + return (extract_batch(pil_images, model, output_type, layers, intermediate_norm),) + + +class KaloscopeFeatureAnalysisNode: + @classmethod + def INPUT_TYPES(cls): + return { + 'required': {'features': ('TENSOR',), 'chart_type': (list(CHART_TYPES),)}, + 'optional': { + 'metric': (['cosine', 'euclidean', 'manhattan'], {'default': 'cosine'}), + 'normalize': ('BOOLEAN', {'default': True, 'tooltip': 'Normalize each image vector to unit L2 norm before distance/clustering.'}), + 'cluster_method': (['kmeans', 'agglomerative', 'dbscan', 'none'], {'default': 'kmeans'}), + 'n_clusters': ('INT', {'default': 3, 'min': 1, 'max': 512}), + 'top_k': ('INT', {'default': 2, 'min': 1, 'max': 511, 'tooltip': 'Neighbor count for relation edges, neighbor ranking and isolation scores.'}), + 'reference_index': ('INT', {'default': 0, 'min': 0, 'max': 511}), + 'tensor_layout': (list(TENSOR_LAYOUTS), {'default': 'auto', 'tooltip': 'For 4D tensors select spatial [B,D,H,W] or layer_tokens [B,L,N,D] explicitly.'}), + 'layer_index': ('INT', {'default': -1, 'min': -128, 'max': 127}), + 'layer_pooling': (['selected', 'mean'], {'default': 'selected'}), + 'token_pooling': (['mean', 'flatten'], {'default': 'mean'}), + 'labels': ('STRING', {'default': '', 'multiline': True, 'tooltip': 'One image name per line or a JSON array, matching feature batch order.'}), + 'images': ('IMAGE', {'tooltip': 'Optional thumbnails only; never used for model inference.'}), + 'seed': ('INT', {'default': 42, 'min': 0, 'max': 2147483647}), + 'perplexity': ('FLOAT', {'default': 5.0, 'min': 0.5, 'max': 100.0}), + 'dbscan_eps': ('FLOAT', {'default': 0.35, 'min': 0.001, 'max': 100.0}), + 'dbscan_min_samples': ('INT', {'default': 2, 'min': 1, 'max': 512}), + 'max_dimensions': ('INT', {'default': 32, 'min': 1, 'max': 128}), + 'heatmap_order': (['cluster', 'input'], {'default': 'cluster'}), + 'grid_width': ('INT', {'default': 0, 'min': 0, 'max': 4096, 'tooltip': 'Patch grid columns; 0 infers a square grid. Spatial maps preserve their H,W.'}), + 'width': ('INT', {'default': 1400, 'min': 512, 'max': 4096, 'step': 64}), + 'height': ('INT', {'default': 1000, 'min': 512, 'max': 4096, 'step': 64}), + }, + } + + RETURN_TYPES = ('IMAGE', 'STRING', 'TENSOR') + RETURN_NAMES = ('visualization', 'analysis_json', 'distance_matrix') + FUNCTION = 'analyze' + CATEGORY = 'Kaloscope/Analysis' + + def analyze(self, features, chart_type='relationship_graph', **kwargs): + return analyze_features(features, chart_type=chart_type, **kwargs) + + +class KaloscopeImageAnalysisNode: + @classmethod + def INPUT_TYPES(cls): + schema = KaloscopeFeatureAnalysisNode.INPUT_TYPES() + schema['required'].pop('features') + schema['required'] = {'image': ('IMAGE',), 'model': ('KALOSCOPE_MODEL',), **schema['required']} + schema['optional'].pop('images') + schema['optional'].update(KaloscopeExtractFeaturesNode.INPUT_TYPES()['optional']) + return schema + + RETURN_TYPES = ('IMAGE', 'STRING', 'TENSOR', 'TENSOR') + RETURN_NAMES = ('visualization', 'analysis_json', 'features', 'distance_matrix') + FUNCTION = 'analyze' + CATEGORY = 'Kaloscope/Analysis' + + def analyze(self, image, model, chart_type='relationship_graph', output_type='default', layers='-1', + intermediate_norm=True, **kwargs): + features = KaloscopeExtractFeaturesNode().extract(image, model, output_type, layers, intermediate_norm)[0] + images = image if image.ndim == 4 else image.unsqueeze(0) + if kwargs.get('tensor_layout', 'auto') == 'auto': + kwargs['tensor_layout'] = output_layout(output_type) + visualization, report, distances = analyze_features(features, chart_type=chart_type, images=images, **kwargs) + return visualization, report, features, distances + + NODE_CLASS_MAPPINGS = { - "LSNetModelLoader": LSNetModelLoader, - "LSNetArtistInference": LSNetArtistInferenceNode, - "LSNetArtistSimilarity": LSNetArtistSimilarityNode, - "LSNetCommonFeatures": LSNetCommonFeaturesNode, - "LSNetClustering": LSNetClusteringNode, - "LSNetFeatureComparison": LSNetFeatureComparisonNode, - "LSNetArtistImageConnector": LSNetArtistImageConnector + 'KaloscopeModelLoader': KaloscopeModelLoader, + 'KaloscopeArtistInference': KaloscopeArtistInferenceNode, + 'KaloscopeArtistSimilarity': KaloscopeArtistSimilarityNode, + 'KaloscopeCommonFeatures': KaloscopeCommonFeaturesNode, + 'KaloscopeClustering': KaloscopeClusteringNode, + 'KaloscopeFeatureComparison': KaloscopeFeatureComparisonNode, + 'KaloscopeArtistImageConnector': KaloscopeArtistImageConnector, + 'KaloscopeExtractFeatures': KaloscopeExtractFeaturesNode, + 'KaloscopeFeatureAnalysis': KaloscopeFeatureAnalysisNode, + 'KaloscopeImageAnalysis': KaloscopeImageAnalysisNode, } NODE_DISPLAY_NAME_MAPPINGS = { - "LSNetModelLoader": "LSNet Model Loader", - "LSNetArtistInference": "LSNet Artist Inference", - "LSNetArtistSimilarity": "LSNet Artist Similarity", - "LSNetCommonFeatures": "LSNet Common Features", - "LSNetClustering": "LSNet Clustering", - "LSNetFeatureComparison": "LSNet Feature Comparison", - "LSNetArtistImageConnector": "LSNet Image Connector" + 'KaloscopeModelLoader': 'Kaloscope Model Loader', + 'KaloscopeArtistInference': 'Kaloscope Artist Inference', + 'KaloscopeArtistSimilarity': 'Kaloscope Artist Similarity', + 'KaloscopeCommonFeatures': 'Kaloscope Common Features', + 'KaloscopeClustering': 'Kaloscope Clustering', + 'KaloscopeFeatureComparison': 'Kaloscope Feature Comparison', + 'KaloscopeArtistImageConnector': 'Kaloscope Image Connector', + 'KaloscopeExtractFeatures': 'Kaloscope Extract Features', + 'KaloscopeFeatureAnalysis': 'Kaloscope Feature Analysis', + 'KaloscopeImageAnalysis': 'Kaloscope Image Analysis', } diff --git a/analysis_cli.py b/analysis_cli.py new file mode 100644 index 0000000..b6e1ac0 --- /dev/null +++ b/analysis_cli.py @@ -0,0 +1,96 @@ +"""Standalone feature extraction, all charts and feature tools; no ComfyUI required.""" +import argparse +import json +from pathlib import Path + +from PIL import Image + +from backend_lsnet.analysis import (extract_batch, cache_bytes, read_cache, thumbnail_batch, + analyze_cached, save_analysis, feature_tools, output_layout) +from backend_lsnet.analysis_api import AnalysisOptions, values +from feature_analysis import CHART_TYPES, prepare_features +from model_loading import FEATURE_OUTPUTS, load_model_bundle + + +def parser(): + result = argparse.ArgumentParser(description=__doc__) + source = result.add_mutually_exclusive_group(required=True) + source.add_argument('--input', type=Path, help='Image file or directory (extract once)') + source.add_argument('--features', type=Path, help='features.npz cache (no model loaded)') + result.add_argument('--model-dir', type=Path) + result.add_argument('--device', default='cuda') + result.add_argument('--output-type', choices=FEATURE_OUTPUTS, default='default') + result.add_argument('--layers', default='-1') + result.add_argument('--no-intermediate-norm', action='store_true') + result.add_argument('--batch-size', type=int, default=2) + result.add_argument('--output', type=Path, default=Path('outputs')) + result.add_argument('--chart-type', choices=CHART_TYPES, default='relationship_graph') + result.add_argument('--all-charts', action='store_true', help='Requires patch tokens/map for patch_energy') + result.add_argument('--operation', choices=['charts', 'common_features', 'similarity', 'compare_groups'], default='charts') + result.add_argument('--groups', type=Path, help='UTF-8 text, one group name per image') + result.add_argument('--options-json', type=Path, help='JSON object with any analysis parameter') + defaults = values(AnalysisOptions()) + for name, default in defaults.items(): + flag = '--' + name.replace('_', '-') + if name == 'labels': + result.add_argument(flag, default=argparse.SUPPRESS, help='One name per line or JSON array') + elif isinstance(default, bool): + result.add_argument(flag, action=argparse.BooleanOptionalAction, default=argparse.SUPPRESS) + else: + result.add_argument(flag, type=type(default), default=argparse.SUPPRESS) + return result + + +def main(args): + options = values(AnalysisOptions()) + if args.options_json: + supplied = json.loads(args.options_json.read_text(encoding='utf-8')) + if not isinstance(supplied, dict) or set(supplied) - set(options): + raise ValueError('options-json must be an object of known analysis options') + options.update(supplied) + options.update({name: value for name, value in vars(args).items() if name in options}) + images = None + if args.features: + cached = read_cache(args.features) + print('Reusing cached features: no model loading or inference', flush=True) + else: + if not args.model_dir: + raise ValueError('--input requires --model-dir') + extensions = {'.jpg', '.jpeg', '.png', '.bmp', '.tiff', '.webp'} + paths = [args.input] if args.input.is_file() else sorted(p for p in args.input.rglob('*') if p.suffix.lower() in extensions) + if not 1 <= len(paths) <= 512: + raise ValueError('Input must contain 1–512 images') + pictures = [] + for path in paths: + with Image.open(path) as image: + pictures.append(image.copy()) + bundle = load_model_bundle(args.model_dir, device=args.device) + features = extract_batch(pictures, bundle, args.output_type, args.layers, not args.no_intermediate_norm, args.batch_size) + cached = {'features': features, 'labels': [str(path.relative_to(args.input)) if args.input.is_dir() else path.name for path in paths], + 'output_type': args.output_type, 'layers': args.layers} + images = thumbnail_batch(pictures) + del bundle + print(f'Extracted {len(paths)} images once: {list(features.shape)}', flush=True) + args.output.mkdir(parents=True, exist_ok=True) + (args.output / 'features.npz').write_bytes(cache_bytes(cached['features'], cached['labels'], cached['output_type'], cached['layers'])) + if args.operation != 'charts': + layout = output_layout(cached['output_type']) if options['tensor_layout'] == 'auto' else options['tensor_layout'] + groups = args.groups.read_text(encoding='utf-8').splitlines() if args.groups else None + report = feature_tools(cached['features'], args.operation, options['reference_index'], groups, + tensor_layout=layout, layer_index=options['layer_index'], layer_pooling=options['layer_pooling'], token_pooling=options['token_pooling']) + (args.output / f'{args.operation}.json').write_text(json.dumps(report, indent=2, ensure_ascii=False, allow_nan=False), encoding='utf-8') + return + charts = CHART_TYPES if args.all_charts else [args.chart_type] + if args.all_charts: + layout = output_layout(cached['output_type']) if options['tensor_layout'] == 'auto' else options['tensor_layout'] + patches = prepare_features(cached['features'], layout, options['layer_index'], options['layer_pooling'], options['token_pooling'])[1] + if patches is None or cached['output_type'] not in ('patch_tokens', 'patch_map', 'intermediate_patch_tokens', 'intermediate_patch_map'): + raise ValueError('--all-charts includes patch_energy; extract patch_tokens or patch_map first') + for chart in charts: + image, report, distances = analyze_cached(cached, chart, images=images, **options) + files = save_analysis(args.output, chart, image, report, distances) + print(files[0], flush=True) + + +if __name__ == '__main__': + main(parser().parse_args()) diff --git a/api_example/cancel.py b/api_example/cancel.py index ce0534a..128835c 100644 --- a/api_example/cancel.py +++ b/api_example/cancel.py @@ -6,7 +6,7 @@ logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) # API 配置 -API_URL = "http://127.0.0.1:7871/lsnet/v1/cancel" +API_URL = "http://127.0.0.1:7871/kaloscope/v1/cancel" USERNAME = "user" # 替换为你的用户名,如果未启用认证可留空 PASSWORD = "password" # 替换为你的密码,如果未启用认证可留空 @@ -39,4 +39,4 @@ if __name__ == "__main__": result = cancel_inference() print(f"Cancel Result: {result}") except Exception as e: - print(f"Error: {str(e)}") \ No newline at end of file + print(f"Error: {str(e)}") diff --git a/api_example/generate.py b/api_example/generate.py index 6e74000..dc72e6a 100644 --- a/api_example/generate.py +++ b/api_example/generate.py @@ -11,7 +11,7 @@ logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) # API 配置 -API_URL = "http://127.0.0.1:7871/lsnet/v1/infer" +API_URL = "http://127.0.0.1:7871/kaloscope/v1/infer" USERNAME = "user" # 替换为你的用户名,如果未启用认证可留空 PASSWORD = "password" # 替换为你的密码,如果未启用认证可留空 OUTPUT_DIR = "outputs" @@ -83,4 +83,4 @@ if __name__ == "__main__": print(f"Results saved to {output_file}") except Exception as e: - print(f"Error: {str(e)}") \ No newline at end of file + print(f"Error: {str(e)}") diff --git a/backend_lsnet/analysis.py b/backend_lsnet/analysis.py new file mode 100644 index 0000000..df0d673 --- /dev/null +++ b/backend_lsnet/analysis.py @@ -0,0 +1,121 @@ +"""Shared feature workflow for ComfyUI, WebUI, standalone UI, CLI and API.""" +import csv +import io +import json +from pathlib import Path + +import numpy as np +from PIL import Image, ImageOps +import torch + +from feature_analysis import analyze_features, prepare_features +from model_loading import FEATURE_OUTPUTS + + +@torch.inference_mode() +def extract_tensor_batch(encoder, batch, output_type='default', layers='-1', intermediate_norm=True): + if hasattr(encoder, 'extract_tensor'): + return encoder.extract_tensor(batch, output_type, layers, intermediate_norm) + if output_type in ('default', 'backbone'): + return encoder(batch, return_features=True) + raise ValueError(f'{output_type} requires a DINOv3 model; LSNet supports default/backbone') + + +@torch.inference_mode() +def extract_batch(images, bundle, output_type='default', layers='-1', intermediate_norm=True, batch_size=2): + if output_type not in FEATURE_OUTPUTS: + raise ValueError(f'Unknown feature output: {output_type}') + if not images or len(images) > 512 or batch_size < 1: + raise ValueError('Provide 1–512 images and a positive batch size') + encoder = bundle['model'] + outputs = [] + for start in range(0, len(images), batch_size): + batch = torch.stack([bundle['transform'](image) for image in images[start:start + batch_size]]).to(bundle['device']) + result = extract_tensor_batch(encoder, batch, output_type, layers, intermediate_norm) + outputs.append(result.float().cpu().contiguous()) + return torch.cat(outputs) + + +def output_layout(output_type): + if output_type == 'patch_map': + return 'spatial' + if output_type.startswith('intermediate_'): + return 'layer_spatial' if output_type.endswith('patch_map') else ( + 'layer_vectors' if output_type in ('intermediate_cls', 'intermediate_mean', 'intermediate_cls_mean') else 'layer_tokens') + return 'auto' + + +def thumbnail_batch(images): + return torch.stack([torch.from_numpy(np.array(ImageOps.pad(image.convert('RGB'), (240, 200), color='white'))).float() / 255 + for image in images]) + + +def cache_bytes(features, labels, output_type='default', layers='-1'): + stream = io.BytesIO() + np.savez_compressed(stream, features=features.detach().float().cpu().numpy(), + labels=np.asarray(labels, dtype=str), output_type=np.asarray(output_type), layers=np.asarray(layers)) + return stream.getvalue() + + +def read_cache(source): + """Only numeric arrays and string metadata; never unpickle uploaded files.""" + with np.load(source, allow_pickle=False) as data: + array = data['features'] + if array.dtype.kind not in 'fiu' or array.ndim < 2 or not 1 <= array.shape[0] <= 512 or not np.isfinite(array).all(): + raise ValueError('Invalid feature cache') + labels = data['labels'].tolist() + if not isinstance(labels, list) or len(labels) != array.shape[0]: + raise ValueError('Cache labels must match features') + output_type = str(data['output_type']) + if output_type not in FEATURE_OUTPUTS: + raise ValueError('Unknown cached feature output') + return {'features': torch.from_numpy(array.astype(np.float32)), 'labels': labels, + 'output_type': output_type, 'layers': str(data['layers'])} + + +def analyze_cached(cached, chart_type='relationship_graph', images=None, **options): + if options.get('tensor_layout', 'auto') == 'auto': + options['tensor_layout'] = output_layout(cached['output_type']) + if not options.get('labels'): + options['labels'] = cached['labels'] + return analyze_features(cached['features'], chart_type=chart_type, images=images, **options) + + +def feature_tools(features, operation='common_features', reference_index=0, groups=None, **layout_options): + vectors = prepare_features(features, **layout_options)[0] + if operation == 'common_features': + return {'common_features': vectors.mean(0).tolist(), 'sample_count': len(vectors)} + if not 0 <= reference_index < len(vectors): + raise ValueError('reference_index is outside the batch') + targets = vectors if operation == 'similarity' else None + if operation == 'compare_groups': + if groups is None or len(groups) != len(vectors): + raise ValueError('Provide one group name per image') + names = sorted(set(str(group) for group in groups)) + targets = np.stack([vectors[np.asarray([str(group) == name for group in groups])].mean(0) for name in names]) + elif operation != 'similarity': + raise ValueError(f'Unknown feature operation: {operation}') + query = vectors[reference_index] + norms = np.linalg.norm(targets, axis=1) * np.linalg.norm(query) + if np.any(norms <= 1e-12): + raise ValueError('Cosine similarity requires nonzero vectors') + scores = (targets @ query / norms).clip(-1, 1) + result = {'reference_index': reference_index, 'similarities': scores.tolist()} + if operation == 'compare_groups': + result.update(groups=names, best_group=names[int(np.argmax(scores))]) + return result + + +def save_analysis(directory, chart_type, image, report, distances): + directory = Path(directory) + directory.mkdir(parents=True, exist_ok=True) + png, json_path, csv_path = [directory / f'{chart_type}.{extension}' for extension in ('png', 'json', 'csv')] + Image.fromarray((image[0].numpy().clip(0, 1) * 255).round().astype(np.uint8)).save(png) + json_path.write_text(report, encoding='utf-8') + labels = json.loads(report)['labels'] + with csv_path.open('w', newline='', encoding='utf-8-sig') as stream: + writer = csv.writer(stream) + writer.writerow(['image', *labels]) + for label, row in zip(labels, distances.tolist()): + writer.writerow([label, *row]) + return [str(path) for path in (png, json_path, csv_path)] diff --git a/backend_lsnet/analysis_api.py b/backend_lsnet/analysis_api.py new file mode 100644 index 0000000..26d03fc --- /dev/null +++ b/backend_lsnet/analysis_api.py @@ -0,0 +1,130 @@ +"""API schemas and handlers for the shared feature/analysis pipeline.""" +import base64 +from io import BytesIO +from typing import Optional, List + +import numpy as np +from PIL import Image +from pydantic import BaseModel, Field +import torch + +from backend_lsnet.analysis import (extract_batch, cache_bytes, read_cache, thumbnail_batch, + analyze_cached, feature_tools, output_layout) +from backend_lsnet.model_paths import get_checkpoint_path +from model_loading import load_model_bundle + + +class AnalysisOptions(BaseModel): + metric: str = 'cosine' + normalize: bool = True + cluster_method: str = 'kmeans' + n_clusters: int = Field(3, ge=1, le=512) + top_k: int = Field(2, ge=1, le=511) + reference_index: int = Field(0, ge=0, le=511) + tensor_layout: str = 'auto' + layer_index: int = -1 + layer_pooling: str = 'selected' + token_pooling: str = 'mean' + labels: List[str] = Field(default_factory=list) + seed: int = Field(42, ge=0, le=2147483647) + perplexity: float = Field(5., gt=0) + dbscan_eps: float = Field(.35, gt=0) + dbscan_min_samples: int = Field(2, ge=1) + max_dimensions: int = Field(32, ge=1, le=128) + heatmap_order: str = 'cluster' + grid_width: int = Field(0, ge=0) + width: int = Field(1400, ge=512, le=4096) + height: int = Field(1000, ge=512, le=4096) + + +class FeaturesRequest(BaseModel): + input_images: List[str] = Field(..., **({'min_length': 1, 'max_length': 512} if hasattr(BaseModel, 'model_dump') + else {'min_items': 1, 'max_items': 512})) + model_name: str = 'Kaloscope' + device: str = 'cuda' + output_type: str = 'default' + layers: str = '-1' + intermediate_norm: bool = True + batch_size: int = Field(2, ge=1, le=512) + labels: List[str] = Field(default_factory=list) + + +class CachedFeaturesRequest(BaseModel): + cache_base64: Optional[str] = None + features: Optional[list] = None + output_type: str = 'default' + labels: List[str] = Field(default_factory=list) + + +class AnalysisRequest(CachedFeaturesRequest): + image_batch: Optional[FeaturesRequest] = None + thumbnail_images: List[str] = Field(default_factory=list) + chart_type: str = 'relationship_graph' + options: AnalysisOptions = Field(default_factory=AnalysisOptions) + + +class FeatureToolsRequest(CachedFeaturesRequest): + operation: str = 'common_features' + reference_index: int = Field(0, ge=0) + groups: Optional[List[str]] = None + tensor_layout: str = 'auto' + layer_index: int = -1 + layer_pooling: str = 'selected' + token_pooling: str = 'mean' + + +def values(model): + return model.model_dump() if hasattr(model, 'model_dump') else model.dict() + + +def extract_request(req, decode): + images = [decode(value) for value in req.input_images] + bundle = load_model_bundle(checkpoint=get_checkpoint_path(req.model_name), device=req.device) + features = extract_batch(images, bundle, req.output_type, req.layers, req.intermediate_norm, req.batch_size) + labels = req.labels or [f'Image {index + 1:02d}' for index in range(len(images))] + if len(labels) != len(images): + raise ValueError('labels must match the image count') + return {'features': features, 'labels': labels, 'output_type': req.output_type, 'layers': req.layers}, images + + +def cached_request(req): + if (req.cache_base64 is None) == (req.features is None): + raise ValueError('Provide exactly one of cache_base64 or features') + if req.cache_base64 is not None: + return read_cache(BytesIO(base64.b64decode(req.cache_base64, validate=True))) + features = torch.tensor(req.features, dtype=torch.float32) + if features.ndim < 2: + raise ValueError('features must have an image batch dimension') + labels = req.labels or [f'Image {index + 1:02d}' for index in range(len(features))] + if len(labels) != len(features): + raise ValueError('labels must match features') + return {'features': features, 'labels': labels, 'output_type': req.output_type, 'layers': '-1'} + + +def serialized_cache(cached): + return {'cache_base64': base64.b64encode(cache_bytes(cached['features'], cached['labels'], cached['output_type'], cached['layers'])).decode(), + 'shape': list(cached['features'].shape), 'labels': cached['labels'], 'output_type': cached['output_type']} + + +def analysis_request(req, decode): + if req.image_batch is not None: + if req.features is not None or req.cache_base64 is not None: + raise ValueError('image_batch cannot be combined with features/cache_base64') + cached, images = extract_request(req.image_batch, decode) + else: + cached = cached_request(req) + images = [decode(value) for value in req.thumbnail_images] + image, report, distances = analyze_cached(cached, req.chart_type, + images=thumbnail_batch(images) if images else None, **values(req.options)) + png = BytesIO() + Image.fromarray((image[0].numpy().clip(0, 1) * 255).round().astype(np.uint8)).save(png, format='PNG') + import json + return {'image_base64': base64.b64encode(png.getvalue()).decode(), 'analysis': json.loads(report), + 'distance_matrix': distances.tolist(), **serialized_cache(cached)} + + +def tools_request(req): + cached = cached_request(req) + layout = output_layout(cached['output_type']) if req.tensor_layout == 'auto' else req.tensor_layout + return feature_tools(cached['features'], req.operation, req.reference_index, req.groups, + tensor_layout=layout, layer_index=req.layer_index, layer_pooling=req.layer_pooling, token_pooling=req.token_pooling) diff --git a/backend_lsnet/analysis_ui.py b/backend_lsnet/analysis_ui.py new file mode 100644 index 0000000..bb75956 --- /dev/null +++ b/backend_lsnet/analysis_ui.py @@ -0,0 +1,154 @@ +"""The same analysis tab is mounted by WebUI and the independent Gradio app.""" +import json +from functools import wraps +from pathlib import Path +import tempfile + +import gradio as gr +from PIL import Image + +from backend_lsnet.analysis import (extract_batch, thumbnail_batch, cache_bytes, read_cache, + analyze_cached, save_analysis, feature_tools, output_layout) +from backend_lsnet.model_paths import get_available_models, get_checkpoint_path +from feature_analysis import CHART_TYPES, TENSOR_LAYOUTS +from model_loading import FEATURE_OUTPUTS, load_model_bundle + + +def ui_errors(callback): + @wraps(callback) + def wrapped(*args, **kwargs): + try: + return callback(*args, **kwargs) + except (ValueError, FileNotFoundError, TypeError, KeyError) as error: + raise gr.Error(str(error)) from error + return wrapped + + +def file_path(value): + return str(value) if isinstance(value, (str, Path)) else value.name + + +@ui_errors +def extract_uploaded(files, model_name, device, output_type, layers, intermediate_norm, batch_size): + if not files: + raise gr.Error('请先上传图片。') + if len(files) > 512: + raise gr.Error('最多支持 512 张图片。') + images, labels = [], [] + for item in files: + with Image.open(file_path(item)) as image: + images.append(image.copy()) + labels.append(Path(file_path(item)).name) + bundle = load_model_bundle(checkpoint=get_checkpoint_path(model_name), device=device) + features = extract_batch(images, bundle, output_type, layers, intermediate_norm, int(batch_size)) + cached = {'features': features, 'labels': labels, 'output_type': output_type, 'layers': layers, + 'images': thumbnail_batch(images)} + directory = Path(tempfile.mkdtemp(prefix='kaloscope-features-')) + path = directory / 'features.npz' + path.write_bytes(cache_bytes(features, labels, output_type, layers)) + return cached, f'已缓存 {len(images)} 张图片,TENSOR {list(features.shape)}。可反复制图,无需再推理。', str(path) + + +@ui_errors +def import_uploaded_cache(file): + if not file: + raise gr.Error('请选择 features.npz。') + cached = read_cache(file_path(file)) + return cached, f'已导入 {len(cached["labels"])} 张图片的缓存,TENSOR {list(cached["features"].shape)}。未加载模型。' + + +@ui_errors +def plot_uploaded_cache(cached, chart_type, *values): + if cached is None: + raise gr.Error('先提取特征或导入缓存。') + options = dict(zip(ANALYSIS_OPTION_NAMES, values)) + image, report, distances = analyze_cached(cached, chart_type, images=cached.get('images'), **options) + directory = tempfile.mkdtemp(prefix='kaloscope-analysis-') + files = save_analysis(directory, chart_type, image, report, distances) + return image[0].numpy(), report, files + + +@ui_errors +def tools_uploaded_cache(cached, operation, reference_index, groups, tensor_layout, layer_index, layer_pooling, token_pooling): + if cached is None: + raise gr.Error('先提取特征或导入缓存。') + layout = output_layout(cached['output_type']) if tensor_layout == 'auto' else tensor_layout + result = feature_tools(cached['features'], operation, int(reference_index), groups.splitlines() if groups.strip() else None, + tensor_layout=layout, layer_index=int(layer_index), layer_pooling=layer_pooling, token_pooling=token_pooling) + return json.dumps(result, ensure_ascii=False, indent=2, allow_nan=False) + + +ANALYSIS_OPTION_NAMES = ( + 'metric', 'normalize', 'cluster_method', 'n_clusters', 'top_k', 'reference_index', + 'tensor_layout', 'layer_index', 'layer_pooling', 'token_pooling', 'labels', 'seed', + 'perplexity', 'dbscan_eps', 'dbscan_min_samples', 'max_dimensions', 'heatmap_order', + 'grid_width', 'width', 'height', +) + + +def build_analysis_tab(): + """Call within a Blocks/Tabs context; no ComfyUI or WebUI dependencies.""" + with gr.TabItem('Features & Analysis'): + state = gr.State(value=None) + gr.Markdown('### 批量特征与分析\n先提取一次特征,再选择图表反复制图。也可以导入 `.npz` 缓存,无需加载模型。') + with gr.Row(): + with gr.Column(): + inputs = gr.File(label='图片批次(顺序以缓存 labels 为准)', file_count='multiple', file_types=['image']) + available = get_available_models() + model = gr.Dropdown(choices=available, value=(available or [None])[0], label='Model Folder') + refresh = gr.Button('刷新模型列表') + device = gr.Dropdown(['cuda', 'cpu'], value='cuda', label='Device') + output_type = gr.Dropdown(list(FEATURE_OUTPUTS), value='default', label='Feature Output') + layers = gr.Textbox(value='-1', label='Intermediate layers(逗号分隔 block 编号)') + norm = gr.Checkbox(value=True, label='Intermediate LayerNorm') + batch = gr.Number(value=2, precision=0, label='Inference batch size') + extract = gr.Button('提取并缓存特征') + with gr.Column(): + cache_input = gr.File(label='导入 features.npz', file_types=['.npz']) + load = gr.Button('导入缓存(不推理)') + status = gr.Textbox(label='缓存状态', interactive=False) + cache_download = gr.File(label='下载特征缓存') + chart = gr.Dropdown(list(CHART_TYPES), value='relationship_graph', label='Chart Type') + with gr.Accordion('分析参数', open=False): + with gr.Row(): + metric = gr.Dropdown(['cosine', 'euclidean', 'manhattan'], value='cosine', label='Distance metric') + normalize = gr.Checkbox(value=True, label='L2 normalize') + cluster = gr.Dropdown(['kmeans', 'agglomerative', 'dbscan', 'none'], value='kmeans', label='Clustering') + clusters = gr.Number(value=3, precision=0, label='Cluster count') + with gr.Row(): + top_k = gr.Number(value=2, precision=0, label='Neighbor count') + reference = gr.Number(value=0, precision=0, label='Query image index (0-based)') + layout = gr.Dropdown(list(TENSOR_LAYOUTS), value='auto', label='Tensor layout') + layer_index = gr.Number(value=-1, precision=0, label='Input layer position') + with gr.Row(): + layer_pooling = gr.Dropdown(['selected', 'mean'], value='selected', label='Layer pooling') + token_pooling = gr.Dropdown(['mean', 'flatten'], value='mean', label='Token pooling') + seed = gr.Number(value=42, precision=0, label='Seed') + perplexity = gr.Number(value=5.0, label='t-SNE perplexity') + with gr.Row(): + eps = gr.Number(value=.35, label='DBSCAN eps') + min_samples = gr.Number(value=2, precision=0, label='DBSCAN min samples') + dimensions = gr.Number(value=32, precision=0, label='Max dimensions') + order = gr.Dropdown(['cluster', 'input'], value='cluster', label='Heatmap order') + with gr.Row(): + grid_width = gr.Number(value=0, precision=0, label='Patch grid columns (0=auto)') + width = gr.Number(value=1400, precision=0, label='Width (px)') + height = gr.Number(value=1000, precision=0, label='Height (px)') + labels = gr.Textbox(value='', lines=3, label='Labels(每行一个,留空使用文件名)') + plot = gr.Button('从缓存生成图表') + visualization = gr.Image(label='Visualization', interactive=False) + report = gr.Textbox(label='Analysis JSON', lines=12, interactive=False) + downloads = gr.File(label='下载 PNG / JSON / 距离矩阵 CSV', file_count='multiple') + with gr.Accordion('共同特征 / 相似度 / 分组比较', open=False): + operation = gr.Dropdown(['common_features', 'similarity', 'compare_groups'], value='common_features', label='Operation') + groups = gr.Textbox(value='', lines=4, label='Group names(分组比较时,每张图片一个组名)') + run_tools = gr.Button('分析缓存特征') + tools_result = gr.Textbox(label='Feature result JSON', lines=8, interactive=False) + controls = [metric, normalize, cluster, clusters, top_k, reference, layout, layer_index, + layer_pooling, token_pooling, labels, seed, perplexity, eps, min_samples, dimensions, + order, grid_width, width, height] + extract.click(extract_uploaded, [inputs, model, device, output_type, layers, norm, batch], [state, status, cache_download]) + load.click(import_uploaded_cache, [cache_input], [state, status]) + plot.click(plot_uploaded_cache, [state, chart, *controls], [visualization, report, downloads]) + run_tools.click(tools_uploaded_cache, [state, operation, reference, groups, layout, layer_index, layer_pooling, token_pooling], [tools_result]) + refresh.click(lambda: gr.update(choices=get_available_models(), value=(get_available_models() or [None])[0]), [], [model]) diff --git a/backend_lsnet/api.py b/backend_lsnet/api.py index 2bc2e56..753121d 100644 --- a/backend_lsnet/api.py +++ b/backend_lsnet/api.py @@ -16,6 +16,8 @@ from pydantic import BaseModel, Field from PIL import Image import numpy as np from backend_lsnet.inference import process_image_from_pil +from backend_lsnet.analysis_api import (FeaturesRequest, AnalysisRequest, FeatureToolsRequest, + extract_request, serialized_cache, analysis_request, tools_request) try: from modules import shared @@ -30,37 +32,20 @@ except ImportError: logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) -def get_available_checkpoints(model_name): - """Get available checkpoint files for the model""" - models_dir = "models/lsnet" - model_dir = os.path.join(models_dir, model_name) - if os.path.exists(model_dir): - checkpoints = [] - for ext in ['*.pth', '*.ckpt', '*.safetensors']: - checkpoints.extend(glob.glob(os.path.join(model_dir, ext))) - return [os.path.basename(f) for f in checkpoints] - return [] - -def get_available_csv(model_name): - """Get available CSV files for the model""" - models_dir = "models/lsnet" - model_dir = os.path.join(models_dir, model_name) - if os.path.exists(model_dir): - csv_files = glob.glob(os.path.join(model_dir, "*.csv")) - return [os.path.basename(f) for f in csv_files] - return [] - -def get_checkpoint_path(model_name, checkpoint_name): - """Get full checkpoint path""" - models_dir = "models/lsnet" - return os.path.join(models_dir, model_name, checkpoint_name) +from backend_lsnet.model_paths import ( + get_available_checkpoints, get_available_csv, get_checkpoint_path, get_class_csv, +) class InferenceRequest(BaseModel): input_image: str = Field(..., description="Input image as Base64 encoded string") - model_name: str = Field('Kaloscope', description="Model name (subfolder in models/lsnet/)") + model_name: str = Field('Kaloscope', description="Model name (subfolder in models/kaloscope/)") device: str = Field('cuda', description="Device to use") top_k: int = Field(5, ge=1, le=20, description="Number of top predictions") threshold: float = Field(0.0, ge=0.0, le=1.0, description="Probability threshold") + mode: str = 'auto' + output_type: str = 'default' + layers: str = '-1' + intermediate_norm: bool = True class InferenceResponse(BaseModel): results: dict = Field(..., description="Inference results") @@ -70,7 +55,7 @@ class CancelResponse(BaseModel): info: str = Field(..., description="Cancel operation result") class Api: - def __init__(self, app: FastAPI, queue_lock: Lock = None, prefix: str = "/lsnet/v1"): + def __init__(self, app: FastAPI, queue_lock: Lock = None, prefix: str = "/kaloscope/v1"): self.app = app self.queue_lock = queue_lock or Lock() self.prefix = prefix @@ -88,7 +73,7 @@ class Api: methods=["POST"], response_model=InferenceResponse, summary="Perform artist style inference", - description="Classify or cluster an image using LSNet artist model." + description="Classify an image or extract features with LSNet or DINOv3." ) self.add_api_route( "cancel", @@ -98,6 +83,10 @@ class Api: summary="Cancel the current inference task", description="Terminates the ongoing inference task." ) + self.add_api_route('features', self.endpoint_features, methods=['POST'], summary='Extract a batch once; return reusable NPZ cache') + self.add_api_route('analyze', self.endpoint_analyze, methods=['POST'], summary='Render any analysis chart from features/cache or an image batch') + self.add_api_route('feature-tools', self.endpoint_feature_tools, methods=['POST'], summary='Common features, similarity or group comparison without inference') + self.add_api_route('models', self.endpoint_models, methods=['GET'], summary='Available model folders and supported chart/feature types') def auth(self, creds: HTTPBasicCredentials = Depends(HTTPBasic())): if not self.credentials: @@ -112,15 +101,15 @@ class Api: ) def add_api_route(self, path: str, endpoint: Callable, **kwargs): - path = f"{self.prefix}/{path}" if self.prefix else path - if self.credentials: - return self.app.add_api_route(path, endpoint, dependencies=[Depends(self.auth)], **kwargs) - return self.app.add_api_route(path, endpoint, **kwargs) + route = f"{self.prefix}/{path}" if self.prefix else path + dependencies = [Depends(self.auth)] if self.credentials else [] + self.app.add_api_route(route, endpoint, dependencies=dependencies, **kwargs) def decode_base64_image(self, base64_str: str) -> Image.Image: try: img_data = base64.b64decode(base64_str, validate=True) - img = Image.open(BytesIO(img_data)).convert("RGB") + img = Image.open(BytesIO(img_data)) + img.load() return img except base64.binascii.Error: raise HTTPException(400, "Invalid Base64 string format") @@ -145,35 +134,17 @@ class Api: checkpoints = get_available_checkpoints(req.model_name) if not checkpoints: raise HTTPException(400, f"No checkpoints found for model {req.model_name}") - checkpoint_name = checkpoints[0] # use first available - checkpoint = get_checkpoint_path(req.model_name, checkpoint_name) + checkpoint = get_checkpoint_path(req.model_name) if not os.path.exists(checkpoint): raise HTTPException(400, f"Checkpoint not found: {checkpoint}") - # Prepare inference arguments - csv_files = get_available_csv(req.model_name) - class_csv = None - if csv_files: - class_csv = os.path.join("models/lsnet", req.model_name, csv_files[0]) # use first available - - # 自动从config.json读取model类型 - model_dir = os.path.join("models/lsnet", req.model_name) - config_path = os.path.join(model_dir, "config.json") - model_type = 'lsnet_xl_artist' # 默认值 - if os.path.exists(config_path): - try: - with open(config_path, 'r', encoding='utf-8') as f: - config = json.load(f) - if 'model' in config and config['model'] in ['lsnet_t_artist', 'lsnet_s_artist', 'lsnet_b_artist', 'lsnet_l_artist', 'lsnet_xl_artist', 'lsnet_xl_artist_448']: - model_type = config['model'] - logger.info(f"Model type loaded from config: {model_type}") - except Exception as e: - logger.warning(f"Failed to load config.json: {e}") - + class_csv = get_class_csv(req.model_name) infer_args = { - "model": model_type, "checkpoint": checkpoint, - "mode": "classify", # default to classify + "mode": req.mode, + "output_type": req.output_type, + "layers": req.layers, + "intermediate_norm": req.intermediate_norm, "device": req.device, "top_k": req.top_k, "threshold": req.threshold, @@ -184,6 +155,8 @@ class Api: results = await self.run_inference(input_image, **infer_args) return InferenceResponse(results=results, info="Inference completed successfully") + except HTTPException: + raise except Exception as e: logger.error(f"Inference failed: {str(e)}") raise HTTPException(500, f"Inference failed: {str(e)}") @@ -192,8 +165,35 @@ class Api: # For simplicity, just return a message since inference is quick return CancelResponse(info="No active inference to cancel") + async def run_feature_job(self, callback): + def locked_job(): + with self.queue_lock: + return callback() + try: + return await asyncio.get_running_loop().run_in_executor(self.executor, locked_job) + except HTTPException: + raise + except (ValueError, FileNotFoundError, TypeError, KeyError) as error: + raise HTTPException(400, str(error)) from error + + async def endpoint_features(self, req: FeaturesRequest): + return await self.run_feature_job(lambda: serialized_cache(extract_request(req, self.decode_base64_image)[0])) + + async def endpoint_analyze(self, req: AnalysisRequest): + return await self.run_feature_job(lambda: analysis_request(req, self.decode_base64_image)) + + async def endpoint_feature_tools(self, req: FeatureToolsRequest): + return await self.run_feature_job(lambda: tools_request(req)) + + async def endpoint_models(self): + from backend_lsnet.model_paths import get_available_models + from feature_analysis import CHART_TYPES, TENSOR_LAYOUTS + from model_loading import FEATURE_OUTPUTS + return {'models': get_available_models(), 'feature_outputs': FEATURE_OUTPUTS, + 'chart_types': CHART_TYPES, 'tensor_layouts': TENSOR_LAYOUTS} + def on_app_started(demo, app): """Called when the webui app starts""" queue_lock = webui_queue_lock or Lock() api = Api(app, queue_lock) - logger.info("LSNet API routes added to webui") \ No newline at end of file + logger.info("Kaloscope API routes added to webui") diff --git a/backend_lsnet/inference.py b/backend_lsnet/inference.py index f9335ca..cb2ba59 100644 --- a/backend_lsnet/inference.py +++ b/backend_lsnet/inference.py @@ -1,112 +1,25 @@ -import os -import json -import tempfile -from pathlib import Path -import torch +"""Inference adapter shared by the standalone UI and API.""" from PIL import Image -import numpy as np -from inference_artist import ( - get_args_parser, load_checkpoint_state, normalize_state_dict_keys, - resolve_num_classes, resolve_feature_dim, load_model, process_single_image, - load_class_mapping -) -from timm.data import resolve_data_config -from timm.data.transforms_factory import create_transform +from model_loading import load_model_bundle +from inference_artist import classify_image, resolved_mode +from backend_lsnet.analysis import extract_batch -def process_image(image_path, model='lsnet_t_artist', checkpoint='', num_classes=None, feature_dim=None, mode='classify', class_csv=None, device='cuda', top_k=5, threshold=0.0): - """ - Process a single image for artist style inference. - Args: - image_path (str): Path to the input image - model (str): Model architecture - checkpoint (str): Path to model checkpoint - num_classes (int): Number of classes - feature_dim (int): Feature dimension - mode (str): Inference mode ('classify', 'cluster', 'both') - class_csv (str): Path to class mapping CSV - device (str): Device to use - top_k (int): Number of top predictions - threshold (float): Probability threshold +def process_image_from_pil(image, model=None, checkpoint='', num_classes=None, feature_dim=None, + mode='auto', class_csv=None, device='cuda', top_k=5, threshold=0.0, + output_type='default', layers='-1', intermediate_norm=True): + bundle = load_model_bundle(checkpoint=checkpoint, model_name=model, device=device, class_csv=class_csv) + model_obj = bundle['model'] + mode = resolved_mode(model_obj, mode) + tensor = bundle['transform'](image).unsqueeze(0) + result = {} + if mode in ('classify', 'both'): + result['classification'] = classify_image(model_obj, tensor, device, bundle['class_mapping'], top_k, threshold) + if mode in ('cluster', 'both'): + result['features'] = extract_batch([image], bundle, output_type, layers, intermediate_norm)[0].tolist() + return result - Returns: - dict: Inference results - """ - # Create temporary directory for output - with tempfile.TemporaryDirectory() as temp_dir: - output_dir = Path(temp_dir) / "output" - output_dir.mkdir(exist_ok=True) - # Prepare arguments - args = get_args_parser().parse_args([ - '--model', model, - '--checkpoint', checkpoint, - '--input', image_path, - '--output', str(output_dir), - '--device', device, - '--top-k', str(top_k), - '--threshold', str(threshold), - '--mode', mode - ]) - - if num_classes is not None: - args.num_classes = num_classes - if feature_dim is not None: - args.feature_dim = feature_dim - if class_csv is not None: - args.class_csv = class_csv - - # Load checkpoint and state - state_dict = load_checkpoint_state(checkpoint) - state_dict = normalize_state_dict_keys(state_dict) - - # Load class mapping - class_mapping = load_class_mapping(class_csv) if class_csv else None - - # Resolve num_classes - args.num_classes = resolve_num_classes(num_classes, class_mapping, state_dict) - - # Resolve feature_dim - args.feature_dim = resolve_feature_dim(feature_dim, state_dict) - - # Load model - model_obj = load_model(args, state_dict) - - # 根据模型配置动态设置输入大小 - from lsnet_model.lsnet_artist import default_cfgs_artist - if args.model in default_cfgs_artist: - model_cfg = default_cfgs_artist[args.model] - configured_input_size = model_cfg.get('input_size', (3, 224, 224))[1] # 获取高度(假设正方形) - if args.input_size != configured_input_size: - args.input_size = configured_input_size - print(f"Auto-setting input_size to {configured_input_size} for model {args.model}") - - # Prepare transform - config = resolve_data_config({'input_size': (3, args.input_size, args.input_size)}, model=model_obj) - transform = create_transform(**config) - - # Process single image - results = process_single_image(args, model_obj, transform, class_mapping) - - return results - -def process_image_from_pil(image, **kwargs): - """ - Process a PIL image for artist style inference. - - Args: - image (PIL.Image): Input image - **kwargs: Other arguments for process_image - - Returns: - dict: Inference results - """ - with tempfile.NamedTemporaryFile(suffix='.png', delete=False) as temp_file: - image.save(temp_file.name) - try: - return process_image(temp_file.name, **kwargs) - finally: - try: - os.unlink(temp_file.name) - except OSError: - pass # Ignore if file is still in use \ No newline at end of file +def process_image(image_path, **kwargs): + with Image.open(image_path) as image: + return process_image_from_pil(image, **kwargs) diff --git a/backend_lsnet/model_paths.py b/backend_lsnet/model_paths.py new file mode 100644 index 0000000..b3bc7a1 --- /dev/null +++ b/backend_lsnet/model_paths.py @@ -0,0 +1,48 @@ +"""Model discovery shared by the independent UI and API.""" +from pathlib import Path +import os + +from model_loading import CHECKPOINT_EXTENSIONS, find_checkpoint, model_folders + + +def models_root(): + if os.environ.get('KALOSCOPE_MODELS_DIR'): + return Path(os.environ['KALOSCOPE_MODELS_DIR']) + try: + from modules import paths + return Path(paths.models_path) + except (ImportError, AttributeError): + return Path(__file__).resolve().parents[1] / 'models' + + +def get_available_models(): + return sorted(model_folders(models_root())) + + +def get_model_dir(model_name): + folders = model_folders(models_root()) + if model_name not in folders: + raise FileNotFoundError(f"Model folder not found: {model_name}") + return folders[model_name] + + +def get_available_checkpoints(model_name): + try: + directory = get_model_dir(model_name) + except FileNotFoundError: + return [] + return [p.name for p in sorted(directory.iterdir()) if p.suffix.lower() in CHECKPOINT_EXTENSIONS] + + +def get_available_csv(model_name): + return [p.name for p in sorted(get_model_dir(model_name).glob("*.csv"))] + + +def get_checkpoint_path(model_name, checkpoint_name=None): + directory = get_model_dir(model_name) + return str(directory / checkpoint_name) if checkpoint_name else str(find_checkpoint(directory)) + + +def get_class_csv(model_name): + path = get_model_dir(model_name) / "class_mapping.csv" + return str(path) if path.is_file() else None diff --git a/backend_lsnet/ui.py b/backend_lsnet/ui.py index 4e6c4f3..0a19bf5 100644 --- a/backend_lsnet/ui.py +++ b/backend_lsnet/ui.py @@ -1,41 +1,14 @@ import gradio as gr from backend_lsnet.inference import process_image_from_pil +from backend_lsnet.analysis_ui import build_analysis_tab +from model_loading import FEATURE_OUTPUTS import os import json -import glob -def get_available_models(): - """Get available model folders from models/lsnet/""" - models_dir = "models/lsnet" - if os.path.exists(models_dir): - subdirs = [d for d in os.listdir(models_dir) if os.path.isdir(os.path.join(models_dir, d))] - if subdirs: - return subdirs - -def get_available_checkpoints(model_name): - """Get available checkpoint files for the model""" - models_dir = "models/lsnet" - model_dir = os.path.join(models_dir, model_name) - if os.path.exists(model_dir): - checkpoints = [] - for ext in ['*.pth', '*.ckpt', '*.safetensors']: - checkpoints.extend(glob.glob(os.path.join(model_dir, ext))) - return [os.path.basename(f) for f in checkpoints] - return [] - -def get_available_csv(model_name): - """Get available CSV files for the model""" - models_dir = "models/lsnet" - model_dir = os.path.join(models_dir, model_name) - if os.path.exists(model_dir): - csv_files = glob.glob(os.path.join(model_dir, "*.csv")) - return [os.path.basename(f) for f in csv_files] - return [] - -def get_checkpoint_path(model_name, checkpoint_name): - """Get full checkpoint path""" - models_dir = "models/lsnet" - return os.path.join(models_dir, model_name, checkpoint_name) +from backend_lsnet.model_paths import ( + get_available_models, get_available_checkpoints, get_available_csv, + get_checkpoint_path, get_class_csv, +) def create_ui(): css = """ @@ -47,61 +20,48 @@ def create_ui(): } """ - block = gr.Blocks(css=css) + block = gr.Blocks(css=css, analytics_enabled=False) with block: - gr.Markdown('# LSNet Artist Inference') + gr.Markdown('# Kaloscope Artist Inference') with gr.Tabs(): with gr.TabItem("Inference"): with gr.Row(): with gr.Column(): - input_image = gr.Image(sources='upload', type="pil", label="Input Image", height=320, elem_classes="contain-image") + image_kwargs = {'source': 'upload'} if int(gr.__version__.split('.')[0]) < 4 else {'sources': ['upload']} + input_image = gr.Image(**image_kwargs, type="pil", label="Input Image", height=320, elem_classes="contain-image") model = gr.Dropdown( choices=get_available_models(), - label="Model Folder", value='Kaloscope' + label="Model Folder", value=(get_available_models() or [None])[0] ) device = gr.Dropdown(['cuda', 'cpu'], label="Device", value='cuda') top_k = gr.Slider(label="Top K", minimum=1, maximum=20, value=5, step=1) threshold = gr.Slider(label="Threshold", minimum=0.0, maximum=1.0, value=0.0, step=0.01) + mode = gr.Dropdown(['auto', 'classify', 'cluster', 'both'], value='auto', label='Mode') + output_type = gr.Dropdown(list(FEATURE_OUTPUTS), value='default', label='Feature Output') + layers = gr.Textbox(value='-1', label='Intermediate layers') + norm = gr.Checkbox(value=True, label='Intermediate LayerNorm') infer_button = gr.Button(value="Infer") with gr.Column(): tag_string = gr.Textbox(label="Formatted Tags", lines=3, interactive=False) result_json = gr.Textbox(label="JSON Results", lines=15, interactive=False) error_message = gr.Markdown("", visible=False) + build_analysis_tab() - def infer(image, model, device, top_k, threshold): + def infer(image, model, device, top_k, threshold, mode, output_type, layers, norm): if image is None: return "Please upload an image.", "", gr.update(visible=True) checkpoints = get_available_checkpoints(model) if not checkpoints: return f"No checkpoints found for model {model}.", "", gr.update(visible=True) - checkpoint_name = checkpoints[0] # use first available - checkpoint = get_checkpoint_path(model, checkpoint_name) - if not os.path.exists(checkpoint): - return f"Checkpoint not found: {checkpoint}", "", gr.update(visible=True) try: - csv_files = get_available_csv(model) - class_csv = None - if csv_files: - class_csv = os.path.join("models/lsnet", model, csv_files[0]) # use first available - - # 自动从config.json读取model类型 - model_dir = os.path.join("models/lsnet", model) - config_path = os.path.join(model_dir, "config.json") - model_type = 'lsnet_xl_artist' # 默认值 - if os.path.exists(config_path): - try: - with open(config_path, 'r', encoding='utf-8') as f: - config = json.load(f) - if 'model' in config and config['model'] in ['lsnet_t_artist', 'lsnet_s_artist', 'lsnet_b_artist', 'lsnet_l_artist', 'lsnet_xl_artist', 'lsnet_xl_artist_448']: - model_type = config['model'] - print(f"Model type loaded from config: {model_type}") - except Exception as e: - print(f"Warning: Failed to load config.json: {e}") - + checkpoint = get_checkpoint_path(model) + class_csv = get_class_csv(model) kwargs = { - 'model': model_type, 'checkpoint': checkpoint, - 'mode': 'classify', # default to classify + 'mode': mode, + 'output_type': output_type, + 'layers': layers, + 'intermediate_norm': norm, 'device': device, 'top_k': top_k, 'threshold': threshold, @@ -109,15 +69,16 @@ def create_ui(): } results = process_image_from_pil(image, **kwargs) tag_string = ",".join([r['class_name'] for r in results.get('classification', [])]) - json_str = json.dumps({r['class_name']: r['probability'] for r in results.get('classification', [])}, ensure_ascii=False) + output = results + json_str = json.dumps(output, ensure_ascii=False) return tag_string, json_str, gr.update(visible=False) except Exception as e: return str(e), "", gr.update(visible=True) infer_button.click( infer, - inputs=[input_image, model, device, top_k, threshold], + inputs=[input_image, model, device, top_k, threshold, mode, output_type, layers, norm], outputs=[tag_string, result_json, error_message] ) - return block \ No newline at end of file + return block diff --git a/feature_analysis.py b/feature_analysis.py new file mode 100644 index 0000000..e407c6d --- /dev/null +++ b/feature_analysis.py @@ -0,0 +1,591 @@ +"""CPU feature analysis and plotting. No model loading or inference is performed.""" +import json +import math +from threading import RLock + +import matplotlib as mpl +from matplotlib.backends.backend_agg import FigureCanvasAgg +from matplotlib.figure import Figure +from matplotlib.patches import FancyArrowPatch +import numpy as np +from PIL import Image +from scipy.cluster.hierarchy import dendrogram, leaves_list, linkage +from scipy.spatial.distance import squareform +from sklearn.cluster import AgglomerativeClustering, DBSCAN, KMeans +from sklearn.decomposition import PCA +from sklearn.manifold import TSNE +from sklearn.metrics import pairwise_distances, silhouette_samples +import torch + +CHART_TYPES = ( + 'relationship_graph', 'distance_heatmap', 'similarity_heatmap', + 'pca_scatter', 'mds_scatter', 'tsne_scatter', 'dendrogram', 'nearest_neighbors', + 'distance_distribution', 'silhouette', 'cluster_sizes', 'pca_variance', + 'feature_statistics', 'feature_heatmap', 'dimension_correlation', + 'cluster_centroid_heatmap', 'outlier_scores', 'patch_energy', +) +TENSOR_LAYOUTS = ('auto', 'vectors', 'tokens', 'spatial', 'layer_vectors', + 'layer_tokens', 'layer_spatial', 'flatten') +COLORS = ('#0072B2', '#D55E00', '#009E73', '#CC79A7', '#E69F00', '#56B4E9', '#666666') +MARKERS = ('o', '^', 's', 'D', 'P', 'X', 'v') +_PLOT_LOCK = RLock() + + +def prepare_features(features, tensor_layout='auto', layer_index=-1, layer_pooling='selected', + token_pooling='mean'): + """Reduce tokens/maps explicitly; first dimension always identifies images.""" + if not isinstance(features, torch.Tensor) or features.ndim < 2: + raise ValueError('features must be a TENSOR with a batch dimension, e.g. [B,D]') + if any(size == 0 for size in features.shape): + raise ValueError('Feature tensor contains an empty dimension') + if features.shape[0] > 512: + raise ValueError('Analysis supports at most 512 images per graph; split the batch explicitly') + data = features.detach().float().cpu().numpy().astype(np.float64) + if not np.isfinite(data).all(): + raise ValueError('Features contain NaN or infinity') + original_shape = list(data.shape) + layout = tensor_layout + if layout == 'auto': + if data.ndim == 4: + raise ValueError('4D tensor is ambiguous: choose spatial [B,D,H,W] or layer_tokens [B,L,N,D]') + layout = {2: 'vectors', 3: 'tokens', 5: 'layer_spatial'}.get(data.ndim) + if layout is None: + raise ValueError('Choose an explicit tensor_layout for this tensor') + if layout not in TENSOR_LAYOUTS: + raise ValueError(f'Unknown tensor layout: {layout}') + ranks = {'vectors': 2, 'tokens': 3, 'spatial': 4, 'layer_vectors': 3, + 'layer_tokens': 4, 'layer_spatial': 5} + if layout != 'flatten' and data.ndim != ranks[layout]: + raise ValueError(f'{layout} requires {ranks[layout]} dimensions; received {data.shape}') + transformations = [f'input layout: {layout}'] + if layout.startswith('layer_'): + if layer_pooling == 'mean': + data = data.mean(axis=1) + transformations.append('mean over selected input layers') + elif layer_pooling == 'selected': + count = data.shape[1] + index = layer_index + count if layer_index < 0 else layer_index + if not 0 <= index < count: + raise ValueError(f'layer_index is outside the input tensor with {count} layers') + data = data[:, index] + transformations.append(f'select input layer position {index}') + else: + raise ValueError('layer_pooling must be selected or mean') + layout = layout.removeprefix('layer_') + patches = None + grid = None + if layout == 'tokens': + patches = data + elif layout == 'spatial': + grid = data.shape[-2:] + patches = data.reshape(data.shape[0], data.shape[1], -1).transpose(0, 2, 1) + if patches is not None: + if token_pooling == 'mean': + vectors = patches.mean(axis=1) + transformations.append('mean over spatial/token positions') + elif token_pooling == 'flatten': + vectors = data.reshape(data.shape[0], -1) + transformations.append('flatten spatial/token positions in original tensor order') + else: + raise ValueError('token_pooling must be mean or flatten') + else: + vectors = data.reshape(data.shape[0], -1) + return vectors, patches, grid, {'input_shape': original_shape, 'tensor_layout': tensor_layout, + 'vector_shape': list(vectors.shape), 'transformations': transformations} + + +def parse_labels(labels, count): + if not labels: + return [f'Image {index + 1:02d}' for index in range(count)] + if isinstance(labels, str): + stripped = labels.strip() + labels = json.loads(stripped) if stripped.startswith('[') else stripped.splitlines() + if not isinstance(labels, (tuple, list)) or len(labels) != count: + raise ValueError(f'labels must contain exactly {count} names, in feature batch order') + return [str(label) for label in labels] + + +def _thumbnails(images, count): + if images is None: + return None + if not isinstance(images, torch.Tensor) or images.ndim != 4 or images.shape[0] != count or images.shape[-1] not in (3, 4): + raise ValueError('images must be a ComfyUI IMAGE batch [B,H,W,3/4] matching features') + array = images.detach().float().cpu().numpy() + if not np.isfinite(array).all(): + raise ValueError('Thumbnail images contain NaN or infinity') + return [Image.fromarray((sample.clip(0, 1) * 255).astype(np.uint8)).convert('RGB') for sample in array] + + +def _classical_mds(distances): + count = len(distances) + center = np.eye(count) - np.ones((count, count)) / count + values, vectors = np.linalg.eigh(-0.5 * center @ (distances ** 2) @ center) + order = np.argsort(values)[::-1] + positive = np.maximum(values[order[:2]], 0) + coordinates = vectors[:, order[:2]] * np.sqrt(positive) + coordinates = np.pad(coordinates, ((0, 0), (0, max(0, 2 - coordinates.shape[1])))) + approximation = pairwise_distances(coordinates) + denominator = np.square(distances).sum() + stress = math.sqrt(np.square(approximation - distances).sum() / denominator) if denominator > 0 else 0.0 + return coordinates, {'method': 'classical MDS', 'normalized_distance_error': stress, + 'negative_eigenvalue_mass': float(np.abs(values[values < -1e-10]).sum())} + + +def _pca(vectors): + if len(vectors) < 2 or not np.any(np.std(vectors, axis=0) > 1e-12): + return np.zeros((len(vectors), 2)), np.zeros(min(vectors.shape)) + estimator = PCA(n_components=min(vectors.shape), svd_solver='full') + result = estimator.fit_transform(vectors) + coordinates = np.pad(result[:, :2], ((0, 0), (0, max(0, 2 - result.shape[1])))) + return coordinates, estimator.explained_variance_ratio_ + + +def _clusters(vectors, distances, method, n_clusters, seed, eps, min_samples, warnings): + count = len(vectors) + if method == 'none': + return np.zeros(count, dtype=int), {'method': 'none', 'n_clusters': 1} + if method == 'dbscan': + groups = DBSCAN(eps=eps, min_samples=min_samples, metric='precomputed').fit_predict(distances) + return groups, {'method': 'DBSCAN', 'metric': 'selected pairwise distance', 'eps': eps, + 'min_samples': min_samples, 'noise_label': -1} + distinct = len(np.unique(vectors, axis=0)) + actual = min(n_clusters, count, distinct) + if actual != n_clusters: + warnings.append(f'n_clusters reduced from {n_clusters} to {actual}: only {count} samples / {distinct} distinct vectors') + if actual == 1: + return np.zeros(count, dtype=int), {'method': method, 'n_clusters': 1} + if method == 'kmeans': + groups = KMeans(n_clusters=actual, random_state=seed, n_init=10).fit_predict(vectors) + details = {'method': 'KMeans', 'n_clusters': actual, 'objective_metric': 'euclidean', + 'note': 'KMeans uses Euclidean vectors even when graph distances use cosine/manhattan'} + elif method == 'agglomerative': + # Use sklearn to cut to exactly the requested cluster count. + groups = AgglomerativeClustering(n_clusters=actual, metric='precomputed', linkage='average').fit_predict(distances) + details = {'method': 'agglomerative', 'n_clusters': actual, 'metric': 'selected distance', 'linkage': 'average'} + else: + raise ValueError(f'Unknown clustering method: {method}') + return groups, details + + +def _short(label): + return label if len(label) <= 26 else label[:23] + '...' + + +def _color(group): + return '#777777' if group < 0 else COLORS[group % len(COLORS)] + + +def _gallery(fig, grid, thumbnails, labels): + cells = grid.subgridspec(math.ceil(len(thumbnails) / 2), 2, hspace=0.25, wspace=0.12) + for index, thumbnail in enumerate(thumbnails): + axis = fig.add_subplot(cells[index // 2, index % 2]) + preview = thumbnail.copy() + preview.thumbnail((180, 150)) + axis.imshow(preview) + title = _short(labels[index]) + if not title.startswith(f'{index + 1:02d}'): + title = f'{index + 1:02d} ' + title + axis.set_title(title, fontsize=8, pad=3) + axis.axis('off') + + +def _main_axes(fig, thumbnails=None, labels=None): + if thumbnails is not None and len(thumbnails) <= 16: + grid = fig.add_gridspec(1, 2, width_ratios=[3.6, 1.4], wspace=0.12) + _gallery(fig, grid[1], thumbnails, labels) + return fig.add_subplot(grid[0]) + return fig.add_subplot(111) + + +def _style_axis(axis): + for side in ('top', 'right'): + axis.spines[side].set_visible(False) + axis.grid(alpha=0.15) + axis.set_axisbelow(True) + + +def _scatter(axis, coordinates, groups, labels): + for group in sorted(set(groups)): + selected = groups == group + axis.scatter(coordinates[selected, 0], coordinates[selected, 1], s=110, + color=_color(group), marker=MARKERS[group % len(MARKERS)], + edgecolors='white', linewidths=1.0, label='Noise' if group < 0 else f'Cluster {group + 1}', zorder=3) + for index, coordinate in enumerate(coordinates): + axis.annotate(_short(labels[index]), coordinate, xytext=(7, 7), textcoords='offset points', fontsize=9, + bbox={'facecolor': 'white', 'alpha': 0.8, 'edgecolor': 'none', 'pad': 1}, zorder=4) + axis.legend(loc='best', frameon=False, fontsize=9) + axis.margins(0.25) + axis.set_aspect('equal', adjustable='datalim') + _style_axis(axis) + + +def _matrix(fig, axis, values, labels, label, cmap, vmin=None, vmax=None): + image = axis.imshow(values, cmap=cmap, interpolation='nearest', vmin=vmin, vmax=vmax, aspect='auto') + axis.set_xticks(range(len(labels)), [_short(name) for name in labels], rotation=45, ha='right', fontsize=9) + axis.set_yticks(range(len(labels)), [_short(name) for name in labels], fontsize=9) + fig.colorbar(image, ax=axis, shrink=0.8, label=label) + if len(labels) <= 14: + lower, upper = image.get_clim() + for row in range(len(labels)): + for col in range(len(labels)): + value = values[row, col] + text_color = 'white' if (value - lower) / (upper - lower or 1) < 0.45 else '#152838' + axis.text(col, row, f'{value:.2f}', ha='center', va='center', fontsize=8, color=text_color) + + +def create_analysis(features, chart_type='relationship_graph', metric='cosine', normalize=True, + cluster_method='kmeans', n_clusters=3, top_k=2, reference_index=0, + tensor_layout='auto', layer_index=-1, layer_pooling='selected', token_pooling='mean', + labels='', images=None, seed=42, perplexity=5.0, dbscan_eps=0.35, dbscan_min_samples=2, + max_dimensions=32, heatmap_order='cluster', grid_width=0, width=1400, height=1000): + """Return a Figure, JSON-safe numeric report, and the original-space distance TENSOR.""" + if chart_type not in CHART_TYPES: + raise ValueError(f'Unknown chart type: {chart_type}') + if metric not in ('cosine', 'euclidean', 'manhattan'): + raise ValueError(f'Unknown distance metric: {metric}') + if top_k < 1 or n_clusters < 1 or max_dimensions < 1 or perplexity <= 0 or dbscan_eps <= 0 or dbscan_min_samples < 1: + raise ValueError('Cluster count, top_k, dimensions, perplexity and DBSCAN parameters must be positive') + if width < 512 or height < 512 or width > 4096 or height > 4096: + raise ValueError('Plot dimensions must be between 512 and 4096 pixels') + raw, patches, grid, preprocessing = prepare_features(features, tensor_layout, layer_index, layer_pooling, token_pooling) + count = len(raw) + if not 0 <= reference_index < count: + raise ValueError(f'reference_index must be between 0 and {count - 1}') + names = parse_labels(labels, count) + thumbnails = _thumbnails(images, count) + warnings = [] + norms = np.linalg.norm(raw, axis=1) + if metric == 'cosine' and np.any(norms <= 1e-12): + raise ValueError('Cosine distance is undefined for zero-norm features; choose Euclidean or fix the inputs') + vectors = raw / np.maximum(norms[:, None], 1e-12) if normalize else raw.copy() + distances = pairwise_distances(vectors, metric=metric) + distances = np.maximum((distances + distances.T) / 2, 0) + np.fill_diagonal(distances, 0) + unit = raw / np.maximum(norms[:, None], 1e-12) + similarities = (unit @ unit.T).clip(-1, 1) + valid_cosine = np.outer(norms > 1e-12, norms > 1e-12) + groups, clustering = _clusters(vectors, distances, cluster_method, n_clusters, seed, dbscan_eps, dbscan_min_samples, warnings) + pca_coordinates, variance = _pca(vectors) + mds_coordinates, mds_details = _classical_mds(distances) + k = min(top_k, max(0, count - 1)) + neighbors = [] + for index in range(count): + order = sorted((other for other in range(count) if other != index), key=lambda other: (distances[index, other], other))[:k] + neighbors.append([{'index': other, 'label': names[other], 'distance': float(distances[index, other]), + 'cosine_similarity': float(similarities[index, other]) if valid_cosine[index, other] else None} for other in order]) + offdiag = distances[np.triu_indices(count, 1)] + centroid = vectors.mean(0) + centered = vectors - centroid + singular = np.linalg.svd(centered, compute_uv=False) + mass = singular ** 2 + probabilities = mass / mass.sum() if mass.sum() > 1e-20 else np.zeros_like(mass) + effective_rank = math.exp(-np.sum(probabilities[probabilities > 0] * np.log(probabilities[probabilities > 0]))) if np.any(probabilities) else 0.0 + ordering = np.arange(count) + tree = linkage(squareform(distances, checks=False), method='average', optimal_ordering=True) if count > 1 else None + if heatmap_order == 'cluster' and tree is not None: + ordering = leaves_list(tree) + elif heatmap_order not in ('input', 'cluster'): + raise ValueError('heatmap_order must be input or cluster') + non_noise = groups >= 0 + silhouette = None + selected_groups = groups[non_noise] + if 1 < len(set(selected_groups)) < len(selected_groups): + silhouette = np.full(count, np.nan) + silhouette[non_noise] = silhouette_samples(distances[np.ix_(non_noise, non_noise)], selected_groups, metric='precomputed') + selected_dimensions = np.argsort(np.var(vectors, axis=0), kind='stable')[::-1][:max_dimensions] + selected_dimensions.sort() + report = { + 'chart_type': chart_type, 'sample_count': count, 'labels': names, 'preprocessing': preprocessing, + 'distance_metric': metric, 'normalize_vectors': bool(normalize), 'normalization': 'row L2' if normalize else 'none', + 'seed': seed, 'clustering': clustering, 'cluster_labels': groups.tolist(), + 'distances': distances.tolist(), 'cosine_similarities': [[float(similarities[i,j]) if valid_cosine[i,j] else None + for j in range(count)] for i in range(count)], 'nearest_neighbors': neighbors, + 'summary': {'mean_pairwise_distance': float(offdiag.mean()) if offdiag.size else None, + 'min_pairwise_distance': float(offdiag.min()) if offdiag.size else None, + 'max_pairwise_distance': float(offdiag.max()) if offdiag.size else None, + 'effective_rank_centered': effective_rank, + 'effective_rank_definition': 'exp(entropy(normalized squared singular values of centered vectors))', + 'raw_feature_norms': norms.tolist()}, + 'pca_explained_variance_ratio': variance.tolist(), 'mds': mds_details, + 'silhouette_scores': [float(value) if np.isfinite(value) else None for value in silhouette] if silhouette is not None else None, + 'selected_dimension_indices': selected_dimensions.tolist(), 'heatmap_sample_order': ordering.tolist(), 'warnings': warnings, + } + plot_data = dict(raw=raw, vectors=vectors, patches=patches, grid=grid, names=names, thumbnails=thumbnails, + distances=distances, similarities=similarities, groups=groups, pca=pca_coordinates, mds=mds_coordinates, + variance=variance, ordering=ordering, tree=tree, silhouette=silhouette, dimensions=selected_dimensions, + neighbors=neighbors, offdiag=offdiag, norms=norms, valid_cosine=valid_cosine) + with _PLOT_LOCK, mpl.rc_context({'font.family': 'sans-serif', 'font.sans-serif': ['DejaVu Sans', 'Microsoft YaHei', 'Noto Sans CJK SC'], + 'font.size': 11, 'axes.spines.top': False, 'axes.spines.right': False, + 'axes.labelcolor': '#243746', 'text.color': '#243746', 'figure.facecolor': 'white', + 'axes.facecolor': 'white', 'savefig.facecolor': 'white'}): + fig = Figure(figsize=(width / 120, height / 120), dpi=120, layout='constrained') + FigureCanvasAgg(fig) + _draw_chart(fig, chart_type, plot_data, report, reference_index, seed, perplexity, grid_width) + subtitle = f'{count} images | {metric} distance | ' + ('L2-normalized vectors' if normalize else 'raw vectors') + if chart_type in ('feature_statistics', 'patch_energy'): + subtitle = f'{count} images | raw input statistics (before vector normalization)' + fig.suptitle(chart_type.replace('_', ' ').title(), fontsize=20, fontweight='bold', x=0.03, ha='left') + fig.supxlabel(subtitle, fontsize=10, color='#536575') + fig.canvas.draw() + return fig, report, torch.from_numpy(distances.astype(np.float32)) + + +def figure_tensor(fig): + """Render a figure directly into a ComfyUI IMAGE without a file roundtrip.""" + with _PLOT_LOCK: + if fig.stale or not hasattr(fig.canvas, 'renderer'): + fig.canvas.draw() + rgb = np.asarray(fig.canvas.buffer_rgba())[:, :, :3].copy() + return torch.from_numpy(rgb).float().unsqueeze(0) / 255 + + +def analyze_features(features, **kwargs): + fig, report, distances = create_analysis(features, **kwargs) + image = figure_tensor(fig) + fig.clear() + return image, json.dumps(report, ensure_ascii=False, allow_nan=False), distances + + +def _draw_chart(fig, chart, data, report, reference, seed, perplexity, grid_width): + names, groups = data['names'], data['groups'] + count = len(names) + distances = data['distances'] + ordered = data['ordering'] + if chart in ('relationship_graph', 'pca_scatter', 'mds_scatter', 'tsne_scatter'): + axis = _main_axes(fig, data['thumbnails'], names) + if chart == 'pca_scatter': + coordinates = data['pca'] + ratio = np.pad(data['variance'], (0, max(0, 2 - len(data['variance'])))) + axis.set(xlabel=f'PC1 ({ratio[0]:.1%} variance)', ylabel=f'PC2 ({ratio[1]:.1%} variance)') + axis.set_title('PCA projection; cluster labels computed in original feature space', fontsize=11) + details = {'method': 'PCA', 'explained_variance_ratio_2d': ratio[:2].tolist()} + elif chart == 'tsne_scatter': + if count < 3: + raise ValueError('t-SNE requires at least 3 images for this visualization') + actual_perplexity = min(perplexity, count - 1) + if actual_perplexity != perplexity: + report['warnings'].append(f't-SNE perplexity reduced to {actual_perplexity} for {count} samples') + if not distances.any(): + coordinates = np.zeros((count, 2)) + kl = 0.0 + report['warnings'].append('All feature distances are zero; t-SNE represented as coincident points') + else: + estimator = TSNE(n_components=2, perplexity=actual_perplexity, metric='precomputed', + init='random', learning_rate='auto', random_state=seed) + coordinates = estimator.fit_transform(distances) + kl = float(estimator.kl_divergence_) + details = {'method': 't-SNE', 'perplexity': actual_perplexity, 'kl_divergence': kl} + axis.set(xlabel='t-SNE axis 1 (arbitrary units)', ylabel='t-SNE axis 2 (arbitrary units)') + axis.set_title('Neighborhood visualization; 2D distances and cluster gaps are not original distances', fontsize=10) + else: + coordinates = data['mds'] + details = report['mds'] + axis.set(xlabel='MDS axis 1', ylabel='MDS axis 2') + error = details['normalized_distance_error'] + axis.set_title(f'Approximate distance layout | relative distance error = {error:.3f}', fontsize=11) + report['projection'] = {**details, 'coordinates': coordinates.tolist()} + if chart == 'relationship_graph': + edges = sorted({tuple(sorted((index, neighbor['index']))) + for index, row in enumerate(data['neighbors']) for neighbor in row}) + report['edges'] = [{'source': left, 'target': right, 'distance': float(distances[left, right])} + for left, right in edges] + for left, right in edges: + segment = coordinates[[left, right]] + within = groups[left] == groups[right] and groups[left] >= 0 + delta = segment[1]-segment[0] + curvature = 0.0 + if within and abs(delta[0]) > 5*abs(delta[1]): + members = sorted(np.flatnonzero(groups==groups[left]),key=lambda i:coordinates[i,0]) + if abs(members.index(left)-members.index(right)) > 1: + curvature = 0.65 + edge = FancyArrowPatch(segment[0], segment[1], arrowstyle='-', connectionstyle=f'arc3,rad={curvature}', + color=_color(groups[left]) if within else '#98A4AD', linewidth=1.6, alpha=0.65, + linestyle='-' if within else '--', zorder=1) + axis.add_patch(edge) + if count <= 16: + midpoint = segment.mean(axis=0) + curvature/2*np.array([delta[1],-delta[0]]) + axis.text(*midpoint, f'{distances[left,right]:.3f}', fontsize=8, color='#536575', + ha='center', bbox={'facecolor': 'white', 'edgecolor': 'none', 'alpha': 0.85, 'pad': 1}) + axis.set_title(axis.get_title() + '\nUnion kNN graph; edge numbers are original-space distances', fontsize=10) + plot_labels = [f'{index+1:02d}' for index in range(count)] if data['thumbnails'] is not None and count<=16 else names + _scatter(axis, coordinates, groups, plot_labels) + elif chart in ('distance_heatmap', 'similarity_heatmap'): + axis = _main_axes(fig) + if chart == 'similarity_heatmap': + if not data['valid_cosine'].all(): + raise ValueError('Cosine similarity heatmap requires nonzero feature vectors') + values = data['similarities'][np.ix_(ordered, ordered)] + _matrix(fig, axis, values, [names[i] for i in ordered], 'Cosine similarity', 'coolwarm', -1, 1) + else: + values = distances[np.ix_(ordered, ordered)] + _matrix(fig, axis, values, [names[i] for i in ordered], report['distance_metric'] + ' distance', 'viridis_r', 0) + axis.set_title('Average-linkage order' if report['heatmap_sample_order'] != list(range(count)) else 'Input order', fontsize=11) + elif chart == 'dendrogram': + axis = _main_axes(fig) + if data['tree'] is None: + axis.text(0.5, 0.5, 'At least 2 images are required for a dendrogram', ha='center', transform=axis.transAxes) + axis.axis('off') + else: + dendrogram(data['tree'], labels=[_short(name) for name in names], ax=axis, + leaf_rotation=35, leaf_font_size=9, above_threshold_color='#0072B2', color_threshold=0) + report['linkage_matrix'] = data['tree'].tolist() + axis.set(xlabel='Images', ylabel=report['distance_metric'] + ' distance') + axis.set_title('Hierarchical relationships | average linkage', fontsize=12) + _style_axis(axis) + elif chart == 'nearest_neighbors': + axis = _main_axes(fig, data['thumbnails'], names) + row = data['neighbors'][reference] + values = [item['distance'] for item in row] + bars = axis.barh(range(len(row)), values, color=[_color(groups[item['index']]) for item in row]) + axis.set_yticks(range(len(row)), [_short(item['label']) for item in row]) + axis.invert_yaxis() + for bar, value in zip(bars, values): + axis.annotate(f'{value:.4f}', (value, bar.get_y() + bar.get_height()/2), xytext=(5, 0), + textcoords='offset points', va='center', fontsize=10) + axis.set_xlim(0, max(values + [1e-9]) * 1.25) + axis.set(xlabel=report['distance_metric'] + ' distance (lower = closer)', ylabel='Nearest images') + axis.set_title('Query: ' + _short(names[reference]), fontsize=12) + report['reference_index'] = reference + _style_axis(axis) + elif chart == 'distance_distribution': + axis = _main_axes(fig) + values = data['offdiag'] + if not len(values): + raise ValueError('Pairwise distance distribution requires at least 2 images') + bins = min(30, max(4, math.ceil(math.sqrt(len(values))))) + counts, edges, _ = axis.hist(values, bins=bins, color='#0072B2', edgecolor='white', alpha=0.9) + for pair, color, title in ((True, '#009E73', 'Within cluster'), (False, '#D55E00', 'Between clusters')): + subset = [distances[i,j] for i in range(count) for j in range(i+1,count) + if groups[i] >= 0 and groups[j] >= 0 and (groups[i] == groups[j]) == pair] + if subset: + axis.axvline(np.mean(subset), color=color, linestyle='--' if pair else ':', + linewidth=2, label=f'{title} mean: {np.mean(subset):.3f}') + report['histogram'] = {'bin_edges': edges.tolist(), 'counts': counts.astype(int).tolist(), + 'pairwise_distances': values.tolist()} + axis.legend(frameon=False) + axis.set(xlabel=report['distance_metric'] + ' distance', ylabel='Number of unordered pairs') + axis.set_title('Each unordered pair counted once; diagonal excluded', fontsize=11) + _style_axis(axis) + elif chart == 'silhouette': + axis = _main_axes(fig) + values = data['silhouette'] + if values is None: + axis.text(0.5, 0.5, 'Silhouette unavailable\nRequires 2 to N-1 non-noise clusters', + transform=axis.transAxes, ha='center', va='center', fontsize=15) + axis.axis('off') + report['warnings'].append('Silhouette unavailable for this clustering') + else: + order = sorted((i for i in range(count) if np.isfinite(values[i])), key=lambda i: (groups[i], values[i])) + axis.barh(range(len(order)), values[order], color=[_color(groups[i]) for i in order]) + axis.set_yticks(range(len(order)), [_short(names[i]) for i in order], fontsize=9) + average = float(np.nanmean(values)) + report['summary']['mean_silhouette'] = average + axis.axvline(average, color='#D55E00', linestyle='--', label=f'Mean: {average:.3f}') + axis.axvline(0, color='#98A4AD', linewidth=1) + axis.set(xlim=(-1,1), xlabel=f'Silhouette coefficient ({report["distance_metric"]})') + axis.legend(frameon=False) + _style_axis(axis) + elif chart == 'cluster_sizes': + axis = _main_axes(fig) + unique, counts = np.unique(groups, return_counts=True) + axis.bar(range(len(unique)), counts, color=[_color(group) for group in unique], edgecolor='white') + axis.set_xticks(range(len(unique)), ['Noise' if group < 0 else f'Cluster {group+1}' for group in unique]) + for index, value in enumerate(counts): + axis.text(index, value, str(value), ha='center', va='bottom', fontsize=12) + report['cluster_sizes'] = {str(int(group)): int(size) for group,size in zip(unique,counts)} + axis.set(ylabel='Image count', ylim=(0, max(counts)*1.2)) + _style_axis(axis) + elif chart == 'pca_variance': + axis = _main_axes(fig) + variance = data['variance'][:32] + indices = np.arange(1, len(variance)+1) + axis.bar(indices, variance, color='#56B4E9', label='Per component') + axis.plot(indices, np.cumsum(variance), color='#D55E00', marker='o', label='Cumulative') + axis.set(xlabel='Principal component', ylabel='Fraction of variance', ylim=(0,1.04)) + axis.set_title(f'Centered effective rank: {report["summary"]["effective_rank_centered"]:.2f}', fontsize=12) + axis.legend(frameon=False) + _style_axis(axis) + elif chart == 'feature_statistics': + axes = fig.subplots(1,3) + statistics = [('Raw L2 norm', data['norms']), ('Mean absolute value', np.abs(data['raw']).mean(1)), + ('Within-vector standard deviation', data['raw'].std(1))] + report['raw_feature_statistics'] = {title: values.tolist() for title,values in statistics} + for axis,(title,values) in zip(axes,statistics): + axis.barh(range(count), values, color=[_color(group) for group in groups]) + axis.set_yticks(range(count), [_short(name) for name in names], fontsize=8) + axis.invert_yaxis() + axis.set(xlabel=title, xlim=(0, max(float(values.max())*1.15, 1e-9))) + _style_axis(axis) + elif chart in ('feature_heatmap', 'dimension_correlation', 'cluster_centroid_heatmap'): + axis = _main_axes(fig) + dimensions = data['dimensions'] + values = data['vectors'][:, dimensions] + if chart == 'dimension_correlation': + centered = values - values.mean(0) + lengths = np.linalg.norm(centered, axis=0) + scale = np.outer(lengths,lengths) + correlation = np.divide(centered.T @ centered, scale, out=np.zeros_like(scale), where=scale>1e-20) + # Undefined correlations (constant columns) stay visibly masked, not treated as zero. + correlation = np.ma.masked_where(scale<=1e-20, correlation.clip(-1,1)) + cmap = mpl.colormaps['coolwarm'].with_extremes(bad='#D9DEE3') + image = axis.imshow(correlation, cmap=cmap, vmin=-1, vmax=1, interpolation='nearest') + axis.set_xticks(range(len(dimensions)), [str(i) for i in dimensions], rotation=90, fontsize=8) + axis.set_yticks(range(len(dimensions)), [str(i) for i in dimensions], fontsize=8) + report['dimension_correlations'] = [[float(correlation[i,j]) if not np.ma.is_masked(correlation[i,j]) else None + for j in range(len(dimensions))] for i in range(len(dimensions))] + axis.set_title('Pearson correlation across images; gray = constant dimension', fontsize=11) + fig.colorbar(image, ax=axis, label='Pearson correlation') + else: + if chart == 'cluster_centroid_heatmap': + unique = [group for group in sorted(set(groups)) if group>=0] + if not unique: + raise ValueError('Cluster centroid heatmap requires at least one non-noise cluster') + values = np.stack([values[groups==group].mean(0) for group in unique]) + row_labels = [f'Cluster {group+1}' for group in unique] + report['cluster_centroid_values'] = values.tolist() + else: + values = values[ordered] + row_labels = [names[i] for i in ordered] + report['displayed_feature_values'] = values.tolist() + limit = max(float(np.abs(values).max()),1e-12) + image = axis.imshow(values, cmap='coolwarm', interpolation='nearest', aspect='auto', vmin=-limit,vmax=limit) + axis.set_xticks(range(len(dimensions)), [str(i) for i in dimensions], rotation=90, fontsize=8) + axis.set_yticks(range(len(row_labels)), [_short(name) for name in row_labels], fontsize=9) + fig.colorbar(image, ax=axis, label='Feature value after selected vector normalization') + axis.set_title('Dimensions with highest across-image variance (ordered by dimension index)', fontsize=11) + axis.set(xlabel='Feature dimension index') + elif chart == 'outlier_scores': + axis = _main_axes(fig) + scores = np.array([np.mean([item['distance'] for item in row]) if row else 0.0 for row in data['neighbors']]) + order = np.argsort(scores)[::-1] + axis.barh(range(count), scores[order], color=[_color(groups[i]) for i in order]) + axis.set_yticks(range(count), [_short(names[i]) for i in order], fontsize=9) + axis.invert_yaxis() + axis.set(xlabel='Mean distance to selected k nearest images (higher = more isolated)') + axis.set_title('Exploratory isolation score; not a calibrated anomaly probability', fontsize=11) + report['isolation_scores'] = scores.tolist() + _style_axis(axis) + elif chart == 'patch_energy': + patches, grid = data['patches'], data['grid'] + if patches is None: + raise ValueError('patch_energy requires patch_tokens/patch_map with tokens or spatial tensor_layout') + if grid is None: + n_tokens = patches.shape[1] + columns = grid_width or math.isqrt(n_tokens) + if columns < 1 or n_tokens % columns or (not grid_width and columns**2 != n_tokens): + raise ValueError('Token count is not a square grid; specify grid_width for patch_energy') + grid = (n_tokens//columns,columns) + energy = np.linalg.norm(patches,axis=-1).reshape(count,*grid) + report['patch_grid_shape'] = list(grid) + report['patch_l2_norms'] = energy.tolist() + columns = min(4,count) + axes = fig.subplots(math.ceil(count/columns),columns,squeeze=False) + lower,upper = float(energy.min()),float(energy.max()) + for index,axis in enumerate(axes.flat): + if index>=count: + axis.axis('off') + continue + image = axis.imshow(energy[index],cmap='viridis',vmin=lower,vmax=upper,interpolation='nearest') + axis.set_title(_short(names[index]),fontsize=10) + axis.set(xlabel='Patch column',ylabel='Patch row') + fig.colorbar(image,ax=list(axes.flat),shrink=0.75,label='Raw patch feature L2 norm (not attention / importance)') diff --git a/inference_artist.py b/inference_artist.py index 72e61ac..8798121 100644 --- a/inference_artist.py +++ b/inference_artist.py @@ -5,7 +5,6 @@ 2. 分类模式:直接输出分类结果 """ import argparse -import csv import json from pathlib import Path from typing import Dict, Optional @@ -14,20 +13,17 @@ import numpy as np import torch import torch.nn.functional as F from PIL import Image -from timm.data import resolve_data_config -from timm.data.transforms_factory import create_transform -from timm.models import create_model - -# Ensure custom LSNet artist models are registered with timm -from lsnet_model import lsnet_artist # noqa: F401 +from model_loading import ( + FEATURE_OUTPUTS, load_model_bundle, load_checkpoint_state, normalize_state_dict_keys, load_class_mapping, +) +from backend_lsnet.analysis import extract_tensor_batch, cache_bytes def get_args_parser(): parser = argparse.ArgumentParser('Artist Style Inference', add_help=False) # 模型参数 - parser.add_argument('--model', default='lsnet_t_artist', type=str, - choices=['lsnet_t_artist', 'lsnet_s_artist', 'lsnet_b_artist', 'lsnet_l_artist', 'lsnet_xl_artist', 'lsnet_xl_artist_448'], + parser.add_argument('--model', default=None, type=str, help='Model architecture') parser.add_argument('--checkpoint', required=True, type=str, help='Path to model checkpoint') @@ -35,13 +31,16 @@ def get_args_parser(): help='Number of classes. If omitted, will try to infer from checkpoint or CSV mapping.') parser.add_argument('--feature-dim', default=None, type=int, help='Feature dimension') - parser.add_argument('--input-size', default=224, type=int, + parser.add_argument('--input-size', default=None, type=int, help='Input image size') # 推理模式 - parser.add_argument('--mode', default='classify', type=str, - choices=['classify', 'cluster', 'both'], + parser.add_argument('--mode', default='auto', type=str, + choices=['auto', 'classify', 'cluster', 'both'], help='Inference mode: classify (with head), cluster (features only), or both') + parser.add_argument('--output-type', choices=FEATURE_OUTPUTS, default='default') + parser.add_argument('--layers', default='-1', help='Intermediate layer indices, comma separated') + parser.add_argument('--no-intermediate-norm', action='store_true') # 输入输出 parser.add_argument('--input', required=True, type=str, @@ -56,8 +55,6 @@ def get_args_parser(): help='Device to use') parser.add_argument('--batch-size', default=32, type=int, help='Batch size for batch inference') - parser.add_argument('--allow-head-reinit', action='store_true', default=False, - help='Allow re-initializing classification head when checkpoint classes mismatch') parser.add_argument('--top-k', default=5, type=int, help='Number of top predictions to show (default: 5)') parser.add_argument('--threshold', default=0.0, type=float, @@ -66,191 +63,41 @@ def get_args_parser(): return parser -def load_checkpoint_state(checkpoint_path: str): - """加载 checkpoint 并返回模型权重""" - checkpoint = torch.load(checkpoint_path, map_location='cpu', weights_only=False) - if isinstance(checkpoint, dict): - if 'model' in checkpoint: - return checkpoint['model'] - if 'model_ema' in checkpoint: - return checkpoint['model_ema'] - return checkpoint +def resolved_mode(model, mode): + if mode not in ('auto', 'classify', 'cluster', 'both'): + raise ValueError(f'Unknown inference mode: {mode}') + if mode == 'auto': + return 'classify' if model.has_classifier else 'cluster' + if mode in ('classify', 'both') and not model.has_classifier: + raise ValueError('Model has no classification head; use --mode auto or --mode cluster') + return mode -def normalize_state_dict_keys(state_dict): - """移除分布式训练前缀等冗余标记""" - normalized = {} - for key, value in state_dict.items(): - if key.startswith('module.'): - new_key = key[len('module.'):] - else: - new_key = key - normalized[new_key] = value - return normalized - - -def resolve_num_classes(num_classes_arg: Optional[int], - class_mapping: Optional[Dict[int, str]], - state_dict) -> int: - """根据参数、CSV 或 checkpoint 推断类别数""" - # 优先使用CSV中的类别数 - if class_mapping: - csv_classes = len(class_mapping) - if num_classes_arg is not None and num_classes_arg != csv_classes: - print(f"[Warning] 提供的 num_classes={num_classes_arg} 与 CSV 中的类别数 {csv_classes} 不一致,已使用 CSV 的值。") - return csv_classes - - # 如果没有CSV,使用参数 - if num_classes_arg is not None: - return num_classes_arg - - # 最后尝试从权重中解析分类头大小 - for key, value in state_dict.items(): - if key.endswith('head.weight') or key.endswith('head.l.weight'): - return value.shape[0] - - raise ValueError('无法推断 num_classes,请提供 CSV 映射文件或显式指定 num_classes 参数。') - - -def resolve_feature_dim(feature_dim_arg: Optional[int], state_dict) -> int: - """根据参数或 checkpoint 推断特征维度""" - if feature_dim_arg is not None: - return feature_dim_arg - - # 尝试从权重中解析特征维度 - # 查找head.bn.weight的维度,这通常是特征维度 - for key, value in state_dict.items(): - if key.endswith('head.bn.weight'): - return value.shape[0] - - # 如果找不到,尝试查找其他可能的特征维度指示器 - for key, value in state_dict.items(): - if 'head' in key and 'weight' in key and len(value.shape) >= 2: - # 对于线性层,输入维度通常是特征维度 - return value.shape[1] if len(value.shape) > 1 else value.shape[0] - - # 默认值 - print("[Warning] 无法从checkpoint推断特征维度,使用默认值384") - return 384 - - -def load_model(args, state_dict): - """加载模型""" - print(f"Loading model: {args.model}") - state_dict = normalize_state_dict_keys(state_dict) - - model = create_model( - args.model, - pretrained=False, - num_classes=args.num_classes, - feature_dim=args.feature_dim, - ) - - model_state = model.state_dict() - adapted_state = {} - mismatched = {} - - for key, value in state_dict.items(): - if key in model_state: - if model_state[key].shape != value.shape: - mismatched[key] = (model_state[key].shape, value.shape) - continue - adapted_state[key] = value - - classifier_keys = [key for key in mismatched if 'head' in key or 'classifier' in key] - other_mismatched = {key: shapes for key, shapes in mismatched.items() if key not in classifier_keys} - - if other_mismatched: - details = '\n'.join([ - f" - {key}: checkpoint {found} -> model {expected}" - for key, (expected, found) in other_mismatched.items() - ]) - raise RuntimeError( - "以下权重尺寸与当前模型不兼容,且无法自动处理,请检查 checkpoint 或模型配置:\n" + details - ) - - require_strict_head = (args.mode in ['classify', 'both']) and not args.allow_head_reinit - - if classifier_keys and require_strict_head: - details = '\n'.join([ - f" - {key}: checkpoint {mismatched[key][1]} -> model {mismatched[key][0]}" - for key in classifier_keys - ]) - raise RuntimeError( - "分类模式下检测到 checkpoint 分类头与当前 num_classes 不一致,已终止加载以避免随机初始化结果。\n" - "请使用与训练数据一致的 checkpoint,或在确认需要重新初始化分类头时添加 --allow-head-reinit," - "或者切换到 --mode cluster 仅提取特征。\n" + details - ) - - if classifier_keys: - print("[Warning] 分类头权重尺寸与当前 num_classes 不一致,将重新初始化以下权重:") - for key in classifier_keys: - expected, found = mismatched[key] - print(f" - {key}: checkpoint {found} -> model {expected}") - # 冲突键已在上文过滤掉,无需额外处理 - - load_result = model.load_state_dict(adapted_state, strict=False) - - if load_result.missing_keys: - print(f"[Info] Missing keys during load: {load_result.missing_keys}") - if load_result.unexpected_keys: - print(f"[Info] Unexpected keys ignored: {load_result.unexpected_keys}") - - if classifier_keys and args.mode in ['classify', 'both']: - if args.allow_head_reinit: - print("[Warning] 分类模式在 --allow-head-reinit 下运行,分类头为随机初始化;为获得可靠结果请提供匹配的数据集 checkpoint。") - else: - print("[Info] 分类模式未开启或无需分类头,已忽略冲突的分类权重。") - - model.to(args.device) - model.eval() - - print(f"Model loaded from {args.checkpoint}") - return model - - -def load_class_mapping(class_csv_path: Optional[str]) -> Optional[Dict[int, str]]: - """加载 CSV 类别映射,返回 class_id -> name 的字典""" - if not class_csv_path: - return None - - csv_path = Path(class_csv_path) - if not csv_path.exists(): - raise FileNotFoundError(f"Class mapping CSV not found: {csv_path}") - - with csv_path.open('r', encoding='utf-8-sig') as f: - reader = csv.DictReader(f) - if not reader.fieldnames or 'class_id' not in reader.fieldnames or 'class_name' not in reader.fieldnames: - raise ValueError('CSV 必须包含 class_id 和 class_name 两列。') - - mapping: Dict[int, str] = {} - for row in reader: - class_id = int(row['class_id']) - class_name = row['class_name'] - mapping[class_id] = class_name - - if not mapping: - raise ValueError(f"CSV {csv_path} 中未找到任何类别映射。") - - return mapping +def load_model(args, state_dict=None): + bundle = load_model_bundle(checkpoint=args.checkpoint, device=args.device, model_name=args.model, + class_csv=args.class_csv, input_size=args.input_size) + args.mode = resolved_mode(bundle['model'], args.mode) + return bundle['model'] def preprocess_image(image_path, transform): """预处理单张图像""" - image = Image.open(image_path).convert('RGB') - tensor = transform(image) + with Image.open(image_path) as image: + tensor = transform(image) return tensor.unsqueeze(0) def classify_image(model, image_tensor, device, class_mapping: Optional[Dict[int, str]] = None, top_k=5, threshold=0.0): """对图像进行分类""" + if not model.has_classifier: + raise ValueError('This checkpoint has no classification head') with torch.no_grad(): image_tensor = image_tensor.to(device) # 使用分类头 logits = model(image_tensor, return_features=False) # 计算概率 - probs = F.softmax(logits, dim=-1) + probs = F.softmax(logits.float(), dim=-1) # Top-K结果 top_probs, top_indices = torch.topk(probs, k=min(top_k, probs.size(-1)), dim=-1) @@ -272,13 +119,13 @@ def classify_image(model, image_tensor, device, class_mapping: Optional[Dict[int return results -def extract_features(model, image_tensor, device): +def extract_features(model, image_tensor, device, output_type='default', layers='-1', intermediate_norm=True): """提取特征向量用于聚类""" with torch.no_grad(): image_tensor = image_tensor.to(device) # 不使用分类头,直接返回特征 - features = model(image_tensor, return_features=True) - return features.cpu().numpy() + features = extract_tensor_batch(model, image_tensor, output_type, layers, intermediate_norm) + return features.float().cpu().numpy() def process_single_image(args, model, transform, class_mapping: Optional[Dict[int, str]] = None): @@ -307,7 +154,7 @@ def process_single_image(args, model, transform, class_mapping: Optional[Dict[in # 聚类模式(提取特征) if args.mode in ['cluster', 'both']: print("\n[Feature Extraction]") - features = extract_features(model, image_tensor, args.device) + features = extract_features(model, image_tensor, args.device, args.output_type, args.layers, not args.no_intermediate_norm) results['features'] = features[0].tolist() print(f"Feature vector shape: {features.shape}") print(f"Feature vector (first 10 dims): {features[0][:10]}") @@ -324,7 +171,7 @@ def process_directory(args, model, transform, class_mapping: Optional[Dict[int, # 支持的图像格式 image_extensions = {'.jpg', '.jpeg', '.png', '.bmp', '.tiff', '.webp'} - image_paths = [p for p in input_dir.glob('**/*') if p.suffix.lower() in image_extensions] + image_paths = sorted(p for p in input_dir.glob('**/*') if p.suffix.lower() in image_extensions) if not image_paths: print(f"No images found in {input_dir}") @@ -340,10 +187,12 @@ def process_directory(args, model, transform, class_mapping: Optional[Dict[int, # 预处理批次 batch_tensors = [] + valid_paths = [] for path in batch_paths: try: tensor = preprocess_image(path, transform) batch_tensors.append(tensor) + valid_paths.append(path) except Exception as e: print(f"Error processing {path.name}: {e}") continue @@ -359,14 +208,14 @@ def process_directory(args, model, transform, class_mapping: Optional[Dict[int, if args.mode in ['classify', 'both']: logits = model(batch_tensor, return_features=False) - probs = F.softmax(logits, dim=-1) + probs = F.softmax(logits.float(), dim=-1) top_probs, top_indices = torch.topk(probs, k=min(args.top_k, probs.size(-1)), dim=-1) if args.mode in ['cluster', 'both']: - features = model(batch_tensor, return_features=True) + features = extract_tensor_batch(model, batch_tensor, args.output_type, args.layers, not args.no_intermediate_norm) # 保存结果 - for j, path in enumerate(batch_paths): + for j, path in enumerate(valid_paths): if j >= len(batch_tensors): continue @@ -397,7 +246,7 @@ def process_directory(args, model, transform, class_mapping: Optional[Dict[int, if args.mode in ['cluster', 'both']: result['features'] = features[j].cpu().numpy().tolist() - all_results[path.name] = result + all_results[str(path.relative_to(input_dir))] = result print(f"Processed {min(i + args.batch_size, len(image_paths))}/{len(image_paths)} images") @@ -405,38 +254,17 @@ def process_directory(args, model, transform, class_mapping: Optional[Dict[int, def main(args): - # 根据模型配置动态设置输入大小 - from lsnet_model.lsnet_artist import default_cfgs_artist - if args.model in default_cfgs_artist: - model_cfg = default_cfgs_artist[args.model] - configured_input_size = model_cfg.get('input_size', (3, 224, 224))[1] # 获取高度(假设正方形) - if args.input_size != configured_input_size: - args.input_size = configured_input_size - print(f"Auto-setting input_size to {configured_input_size} for model {args.model} (from config)") - - # 创建输出目录 + if args.batch_size < 1 or args.top_k < 1: + raise ValueError('batch-size and top-k must be positive') + bundle = load_model_bundle(checkpoint=args.checkpoint, model_name=args.model, device=args.device, + class_csv=args.class_csv, input_size=args.input_size) + model, transform, class_mapping = bundle['model'], bundle['transform'], bundle['class_mapping'] + args.mode = resolved_mode(model, args.mode) + print(f"Loaded {bundle['model_type']}: classifier={bundle['has_classifier']}, " + f"features={bundle['feature_dim']} ({bundle['feature_source']}), input={bundle['input_size']}") output_dir = Path(args.output) output_dir.mkdir(parents=True, exist_ok=True) - - if args.mode in ['classify', 'both'] and not args.class_csv: - raise ValueError('分类或混合模式下必须提供 --class-csv,且需使用训练阶段导出的映射文件。') - # 加载类别映射 - class_mapping = load_class_mapping(args.class_csv) - - # 加载 checkpoint 并解析类别数 - state_dict = load_checkpoint_state(args.checkpoint) - state_dict = normalize_state_dict_keys(state_dict) - args.num_classes = resolve_num_classes(args.num_classes, class_mapping, state_dict) - args.feature_dim = resolve_feature_dim(args.feature_dim, state_dict) - - # 加载模型 - model = load_model(args, state_dict) - - # 创建数据转换 - config = resolve_data_config({'input_size': (3, args.input_size, args.input_size)}, model=model) - transform = create_transform(**config) - # 判断输入类型 input_path = Path(args.input) @@ -449,6 +277,9 @@ def main(args): with open(output_file, 'w', encoding='utf-8') as f: json.dump(results, f, indent=2, ensure_ascii=False) print(f"\nResults saved to: {output_file}") + if results and 'features' in results: + (output_dir / 'features.npz').write_bytes(cache_bytes(torch.tensor([results['features']]), + [input_path.name], args.output_type, args.layers)) elif input_path.is_dir(): # 目录批量处理 @@ -472,6 +303,8 @@ def main(args): if features_list: features_array = np.array(features_list) np.save(output_dir / "features.npy", features_array) + (output_dir / 'features.npz').write_bytes(cache_bytes(torch.from_numpy(features_array), + image_names, args.output_type, args.layers)) with open(output_dir / "image_names.txt", 'w') as f: f.write('\n'.join(image_names)) print(f"Feature matrix saved: {output_dir / 'features.npy'}") diff --git a/install.py b/install.py index cafd050..b6fc854 100644 --- a/install.py +++ b/install.py @@ -49,8 +49,10 @@ requirements = [ "timm>=1.0.20", "numpy>=1.19.0", "Pillow>=8.0.0", - "scikit-learn>=1.0.0", - "tqdm>=4.60.0" + "scikit-learn>=1.2.0", + "scipy>=1.8.0", + "tqdm>=4.60.0", + "safetensors>=0.4.0" ] # Add triton-windows only on Windows @@ -59,7 +61,7 @@ if sys.platform == 'win32': # Add webui packages webui_packages = [ - "gradio", + "gradio>=3.41,<6", "fastapi", "uvicorn", "python-multipart" @@ -68,4 +70,4 @@ requirements.extend(webui_packages) for req in requirements: if not is_installed(req): - launch.run_pip(f"install {req}", f"sd-webui-lsnet requirement: {req}") + launch.run_pip(f"install {req}", f"comfyui-kaloscope requirement: {req}") diff --git a/kaloscope_dinov3/LICENSE.md b/kaloscope_dinov3/LICENSE.md new file mode 100644 index 0000000..f531b1e --- /dev/null +++ b/kaloscope_dinov3/LICENSE.md @@ -0,0 +1,66 @@ +# DINOv3 License + +*Last Updated: August 19, 2025* + +**“Agreement”** means the terms and conditions for use, reproduction, distribution and modification of the DINO Materials set forth herein. + +**“DINO Materials”** means, collectively, Documentation and the models, software and algorithms, including machine-learning model code, trained model weights, inference-enabling code, training-enabling code, fine-tuning enabling code, and other elements of the foregoing distributed by Meta and made available under this Agreement. + +**“Documentation”** means the specifications, manuals and documentation accompanying +DINO Materials distributed by Meta. + +**“Licensee”** or **“you”** means you, or your employer or any other person or entity (if you are entering into this Agreement on such person or entity’s behalf), of the age required under applicable laws, rules or regulations to provide legal consent and that has legal authority to bind your employer or such other person or entity if you are entering in this Agreement on their behalf. + +**“Meta”** or **“we”** means Meta Platforms Ireland Limited (if you are located in or, if you are an entity, your principal place of business is in the EEA or Switzerland) or Meta Platforms, Inc. (if you are located outside of the EEA or Switzerland). + +**“Sanctions”** means any economic or trade sanctions or restrictions administered or enforced by the United States (including the Office of Foreign Assets Control of the U.S. Department of the Treasury (“OFAC”), the U.S. Department of State and the U.S. Department of Commerce), the United Nations, the European Union, or the United Kingdom. + +**“Trade Controls”** means any of the following: Sanctions and applicable export and import controls. + +By clicking “I Accept” below or by using or distributing any portion or element of the DINO Materials, you agree to be bound by this Agreement. + +## 1. License Rights and Redistribution. + +a. Grant of Rights. You are granted a non-exclusive, worldwide, non-transferable and royalty-free limited license under Meta’s intellectual property or other rights owned by Meta embodied in the DINO Materials to use, reproduce, distribute, copy, create derivative works of, and make modifications to the DINO Materials. + +b. Redistribution and Use. + +i. Distribution of DINO Materials, and any derivative works thereof, are subject to the terms of this Agreement. If you distribute or make the DINO Materials, or any derivative works thereof, available to a third party, you may only do so under the terms of this Agreement and you shall provide a copy of this Agreement with any such DINO Materials. + +ii. If you submit for publication the results of research you perform on, using, or otherwise in connection with DINO Materials, you must acknowledge the use of DINO Materials in your publication. + +iii. Your use of the DINO Materials must comply with applicable laws and regulations, including Trade Control Laws and applicable privacy and data protection laws. + +iv. Your use of the DINO Materials will not involve or encourage others to reverse engineer, decompile or discover the underlying components of the DINO Materials. + +v. You are not the target of Trade Controls and your use of DINO Materials must comply with Trade Controls. You agree not to use, or permit others to use, DINO Materials for any activities subject to the International Traffic in Arms Regulations (ITAR) or end uses prohibited by Trade Controls, including those related to military or warfare purposes, nuclear industries or applications, espionage, or the development or use of guns or illegal weapons. + +## 2. User Support. + +Your use of the DINO Materials is done at your own discretion; Meta does not process any information nor provide any service in relation to such use. Meta is under no obligation to provide any support services for the DINO Materials. Any support provided is “as is”, “with all faults”, and without warranty of any kind. + +## 3. Disclaimer of Warranty. + +UNLESS REQUIRED BY APPLICABLE LAW, THE DINO MATERIALS AND ANY OUTPUT AND RESULTS THEREFROM ARE PROVIDED ON AN “AS IS” BASIS, WITHOUT WARRANTIES OF ANY KIND, AND META DISCLAIMS ALL WARRANTIES OF ANY KIND, BOTH EXPRESS AND IMPLIED, INCLUDING, WITHOUT LIMITATION, ANY WARRANTIES OF TITLE, NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. YOU ARE SOLELY RESPONSIBLE FOR DETERMINING THE APPROPRIATENESS OF USING OR REDISTRIBUTING THE DINO MATERIALS AND ASSUME ANY RISKS ASSOCIATED WITH YOUR USE OF THE DINO MATERIALS AND ANY OUTPUT AND RESULTS. + +## 4. Limitation of Liability. + +IN NO EVENT WILL META OR ITS AFFILIATES BE LIABLE UNDER ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, TORT, NEGLIGENCE, PRODUCTS LIABILITY, OR OTHERWISE, ARISING OUT OF THIS AGREEMENT, FOR ANY LOST PROFITS OR ANY DIRECT OR INDIRECT, SPECIAL, CONSEQUENTIAL, INCIDENTAL, EXEMPLARY OR PUNITIVE DAMAGES, EVEN IF META OR ITS AFFILIATES HAVE BEEN ADVISED OF THE POSSIBILITY OF ANY OF THE FOREGOING. + +## 5. Intellectual Property. + +a. Subject to Meta’s ownership of DINO Materials and derivatives made by or for Meta, with respect to any derivative works and modifications of the DINO Materials that are made by you, as between you and Meta, you are and will be the owner of such derivative works and modifications. + +b. If you institute litigation or other proceedings against Meta or any entity (including a cross-claim or counterclaim in a lawsuit) alleging that the DINO Materials, outputs or results, or any portion of any of the foregoing, constitutes infringement of intellectual property or other rights owned or licensable by you, then any licenses granted to you under this Agreement shall terminate as of the date such litigation or claim is filed or instituted. You will indemnify and hold harmless Meta from and against any claim by any third party arising out of or related to your use or distribution of the DINO Materials. + +## 6. Term and Termination. + +The term of this Agreement will commence upon your acceptance of this Agreement or access to the DINO Materials and will continue in full force and effect until terminated in accordance with the terms and conditions herein. Meta may terminate this Agreement if you are in breach of any term or condition of this Agreement. Upon termination of this Agreement, you shall delete and cease use of the DINO Materials. Sections 3, 4 and 7 shall survive the termination of this Agreement. + +## 7. Governing Law and Jurisdiction. + +This Agreement will be governed and construed under the laws of the State of California without regard to choice of law principles, and the UN Convention on Contracts for the International Sale of Goods does not apply to this Agreement. The courts of California shall have exclusive jurisdiction of any dispute arising out of this Agreement. + +## 8. Modifications and Amendments. + +Meta may modify this Agreement from time to time; provided that they are similar in spirit to the current version of the Agreement, but may differ in detail to address new problems or concerns. All such changes will be effective immediately. Your continued use of the DINO Materials after any modification to this Agreement constitutes your agreement to such modification. Except as provided in this Agreement, no modification or addition to any provision of this Agreement will be binding unless it is in writing and signed by an authorized representative of both you and Meta. diff --git a/kaloscope_dinov3/README.md b/kaloscope_dinov3/README.md new file mode 100644 index 0000000..bfedd40 --- /dev/null +++ b/kaloscope_dinov3/README.md @@ -0,0 +1,13 @@ +# Vendored DINOv3 inference code + +Copied from the sibling kaloscope-dinov3 repository: dinov3 hub, models, layers, +utils. The required backbone factory and deterministic image transforms were +extracted into architecture.py and preprocessing.py. Absolute imports use the private +kaloscope_dinov3 namespace to avoid other ComfyUI extensions' dinov3 packages. +No fine-tuning model, training scope controls, datasets, data augmentation, +optimizers, losses or training entry points are included. Loading is offline +(pretrained=False). Core model/layer definitions retain the original architecture +implementation to preserve checkpoint compatibility. + +The original copyright headers and DINOv3 License Agreement (LICENSE.md) apply +to these files. Other plugin code retains its existing license. diff --git a/kaloscope_dinov3/__init__.py b/kaloscope_dinov3/__init__.py new file mode 100644 index 0000000..b70f615 --- /dev/null +++ b/kaloscope_dinov3/__init__.py @@ -0,0 +1,6 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +__version__ = "0.0.1" diff --git a/kaloscope_dinov3/architecture.py b/kaloscope_dinov3/architecture.py new file mode 100644 index 0000000..acdf815 --- /dev/null +++ b/kaloscope_dinov3/architecture.py @@ -0,0 +1,33 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +"""Construct DINOv3 inference backbones without downloading weights.""" +import inspect + +from .hub import backbones + +MODEL_NAMES = tuple(name for name in vars(backbones) if name.startswith("dinov3_")) + + +def build_backbone(config): + name = config.get("name", "dinov3_vits16") + kwargs = config.get("kwargs", {}) + if name in ("custom_vit", "custom_convnext"): + if name == "custom_vit": + from .models.vision_transformer import DinoVisionTransformer + factory = DinoVisionTransformer + else: + from .models.convnext import ConvNeXt + factory = ConvNeXt + unknown = set(kwargs) - set(inspect.signature(factory).parameters) + if unknown: + raise ValueError(f"Unknown custom architecture options: {sorted(unknown)}") + backbone = factory(**kwargs) + if name == "custom_vit": + backbone.init_weights() + return backbone + if name not in MODEL_NAMES: + raise ValueError(f"Unknown model {name}; available: {MODEL_NAMES}, custom_vit, custom_convnext") + return getattr(backbones, name)(pretrained=False, **kwargs) diff --git a/kaloscope_dinov3/hub/__init__.py b/kaloscope_dinov3/hub/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/kaloscope_dinov3/hub/backbones.py b/kaloscope_dinov3/hub/backbones.py new file mode 100644 index 0000000..b81e8ea --- /dev/null +++ b/kaloscope_dinov3/hub/backbones.py @@ -0,0 +1,612 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +import os +from enum import Enum +from typing import List, Optional, Union +from urllib.parse import urlparse +from pathlib import Path + +from .utils import _DINOV3_BASE_URL, _safe_load_state_dict_from_url + + +class Weights(Enum): + LVD1689M = "LVD1689M" + SAT493M = "SAT493M" + + +def is_url(path: str) -> bool: + parsed = urlparse(path) + return parsed.scheme in ("https", "file") + + +def convert_path_or_url_to_url(path: str) -> str: + if is_url(path): + return path + return Path(path).expanduser().resolve().as_uri() + + +def _make_dinov3_vit_model_arch( + *, + patch_size: int = 16, + compact_arch_name: str = "vitb", +): + if "plus" in compact_arch_name: + model_arch = compact_arch_name.replace("plus", f"{patch_size}plus") + else: + model_arch = f"{compact_arch_name}{patch_size}" + return model_arch + + +def _make_dinov3_vit_model_url( + *, + patch_size: int = 16, + compact_arch_name: str = "vitb", + version: Optional[str] = None, + weights: Union[Weights, str] = Weights.LVD1689M, + hash: Optional[str] = None, +): + model_name = "dinov3" + model_arch = _make_dinov3_vit_model_arch(patch_size=patch_size, compact_arch_name=compact_arch_name) + version_suffix = f"_{version}" if version else "" + weights_name = weights.value.lower() + hash_suffix = f"-{hash}" if hash else "" + model_dir = f"{model_name}_{model_arch}" + model_filename = f"{model_name}_{model_arch}_pretrain_{weights_name}{version_suffix}{hash_suffix}.pth" + return os.path.join(_DINOV3_BASE_URL, model_dir, model_filename) + + +def _make_dinov3_vit( + *, + img_size: int = 224, + patch_size: int = 16, + in_chans: int = 3, + compact_arch_name: str = "vitb", + pos_embed_rope_base: float = 100.0, + pos_embed_rope_min_period: float | None = None, + pos_embed_rope_max_period: float | None = None, + pos_embed_rope_normalize_coords: str = "separate", + pos_embed_rope_shift_coords: float | None = None, + pos_embed_rope_jitter_coords: float | None = None, + pos_embed_rope_rescale_coords: float | None = None, + pos_embed_rope_dtype: str = "fp32", + embed_dim: int = 768, + depth: int = 12, + num_heads: int = 12, + ffn_ratio: float = 4.0, + qkv_bias: bool = True, + drop_path_rate: float = 0.0, + layerscale_init: float | None = None, + norm_layer: str = "layernorm", + ffn_layer: str = "mlp", + ffn_bias: bool = True, + proj_bias: bool = True, + n_storage_tokens: int = 0, + mask_k_bias: bool = False, + pretrained: bool = True, + version: Optional[str] = None, + weights: Union[Weights, str] = Weights.LVD1689M, + hash: Optional[str] = None, + check_hash: bool = False, + **kwargs, +): + from ..models.vision_transformer import DinoVisionTransformer + + vit_kwargs = dict( + img_size=img_size, + patch_size=patch_size, + in_chans=in_chans, + pos_embed_rope_base=pos_embed_rope_base, + pos_embed_rope_min_period=pos_embed_rope_min_period, + pos_embed_rope_max_period=pos_embed_rope_max_period, + pos_embed_rope_normalize_coords=pos_embed_rope_normalize_coords, + pos_embed_rope_shift_coords=pos_embed_rope_shift_coords, + pos_embed_rope_jitter_coords=pos_embed_rope_jitter_coords, + pos_embed_rope_rescale_coords=pos_embed_rope_rescale_coords, + pos_embed_rope_dtype=pos_embed_rope_dtype, + embed_dim=embed_dim, + depth=depth, + num_heads=num_heads, + ffn_ratio=ffn_ratio, + qkv_bias=qkv_bias, + drop_path_rate=drop_path_rate, + layerscale_init=layerscale_init, + norm_layer=norm_layer, + ffn_layer=ffn_layer, + ffn_bias=ffn_bias, + proj_bias=proj_bias, + n_storage_tokens=n_storage_tokens, + mask_k_bias=mask_k_bias, + ) + vit_kwargs.update(**kwargs) + model = DinoVisionTransformer(**vit_kwargs) + if pretrained: + if type(weights) is Weights and weights not in {Weights.LVD1689M, Weights.SAT493M}: + raise ValueError(f"Unsupported weights for the backbone: {weights}") + elif type(weights) is Weights: + url = _make_dinov3_vit_model_url( + patch_size=patch_size, + compact_arch_name=compact_arch_name, + version=version, + weights=weights, + hash=hash, + ) + else: + url = convert_path_or_url_to_url(weights) + state_dict = _safe_load_state_dict_from_url(url, map_location="cpu", check_hash=check_hash) + model.load_state_dict(state_dict, strict=True) + else: + model.init_weights() + return model + + +def _make_dinov3_convnext_model_url( + *, + compact_arch_name: str = "convnext_base", + weights: Union[Weights, str] = Weights.LVD1689M, + hash: Optional[str] = None, +): + model_name = "dinov3" + weights_name = weights.value.lower() + hash_suffix = f"-{hash}" if hash else "" + + model_dir = f"{model_name}_{compact_arch_name}" + model_filename = f"{model_name}_{compact_arch_name}_pretrain_{weights_name}{hash_suffix}.pth" + return os.path.join(_DINOV3_BASE_URL, model_dir, model_filename) + + +def _make_dinov3_convnext( + in_chans: int = 3, + depths: List[int] = [3, 3, 27, 3], + dims: List[int] = [128, 256, 512, 1024], + compact_arch_name: str = "convnext_base", + drop_path_rate: float = 0.0, + layer_scale_init_value: float = 1e-6, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD1689M, + hash: Optional[str] = None, + **kwargs, +): + from ..models.convnext import ConvNeXt + + model_kwargs = dict( + in_chans=in_chans, + depths=depths, + dims=dims, + drop_path_rate=drop_path_rate, + layer_scale_init_value=layer_scale_init_value, + ) + model_kwargs.update(**kwargs) + model = ConvNeXt(**model_kwargs) + if pretrained: + if type(weights) is Weights and weights not in {Weights.LVD1689M, Weights.SAT493M}: + raise ValueError(f"Unsupported weights for the backbone: {weights}") + elif type(weights) is Weights: + url = _make_dinov3_convnext_model_url( + compact_arch_name=compact_arch_name, + weights=weights, + hash=hash, + ) + else: + url = convert_path_or_url_to_url(weights) + state_dict = _safe_load_state_dict_from_url(url, map_location="cpu") + model.load_state_dict(state_dict, strict=True) + return model + + +def dinov3_vits16( + *, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD1689M, + check_hash: bool = False, + **kwargs, +): + if "hash" not in kwargs: + kwargs["hash"] = "08c60483" + kwargs["version"] = None + return _make_dinov3_vit( + img_size=224, + patch_size=16, + in_chans=3, + pos_embed_rope_base=100, + pos_embed_rope_normalize_coords="separate", + pos_embed_rope_rescale_coords=2, + pos_embed_rope_dtype="fp32", + embed_dim=384, + depth=12, + num_heads=6, + ffn_ratio=4, + qkv_bias=True, + drop_path_rate=0.0, + layerscale_init=1.0e-05, + norm_layer="layernormbf16", + ffn_layer="mlp", + ffn_bias=True, + proj_bias=True, + n_storage_tokens=4, + mask_k_bias=True, + pretrained=pretrained, + weights=weights, + compact_arch_name="vits", + check_hash=check_hash, + **kwargs, + ) + + +def dinov3_vits16plus( + *, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD1689M, + check_hash: bool = False, + **kwargs, +): + if "hash" not in kwargs: + kwargs["hash"] = "4057cbaa" + kwargs["version"] = None + return _make_dinov3_vit( + img_size=224, + patch_size=16, + in_chans=3, + pos_embed_rope_base=100, + pos_embed_rope_normalize_coords="separate", + pos_embed_rope_rescale_coords=2, + pos_embed_rope_dtype="fp32", + embed_dim=384, + depth=12, + num_heads=6, + ffn_ratio=6, + qkv_bias=True, + drop_path_rate=0.0, + layerscale_init=1.0e-05, + norm_layer="layernormbf16", + ffn_layer="swiglu", + ffn_bias=True, + proj_bias=True, + n_storage_tokens=4, + mask_k_bias=True, + pretrained=pretrained, + weights=weights, + compact_arch_name="vitsplus", + check_hash=check_hash, + **kwargs, + ) + + +def dinov3_vitb16( + *, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD1689M, + check_hash: bool = False, + **kwargs, +): + if "hash" not in kwargs: + kwargs["hash"] = "73cec8be" + kwargs["version"] = None + return _make_dinov3_vit( + img_size=224, + patch_size=16, + in_chans=3, + pos_embed_rope_base=100, + pos_embed_rope_normalize_coords="separate", + pos_embed_rope_rescale_coords=2, + pos_embed_rope_dtype="fp32", + embed_dim=768, + depth=12, + num_heads=12, + ffn_ratio=4, + qkv_bias=True, + drop_path_rate=0.0, + layerscale_init=1.0e-05, + norm_layer="layernormbf16", + ffn_layer="mlp", + ffn_bias=True, + proj_bias=True, + n_storage_tokens=4, + mask_k_bias=True, + pretrained=pretrained, + weights=weights, + compact_arch_name="vitb", + check_hash=check_hash, + **kwargs, + ) + + +def dinov3_vitl16( + *, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD1689M, + check_hash: bool = False, + **kwargs, +): + untie_global_and_local_cls_norm = False + if weights == Weights.LVD1689M: + if "hash" not in kwargs: + kwargs["hash"] = "8aa4cbdd" + elif weights == Weights.SAT493M: + if "hash" not in kwargs: + kwargs["hash"] = "eadcf0ff" + untie_global_and_local_cls_norm = True + elif type(weights) is str: + import re + + pattern = r"-(.{8}).pth" + matches = re.findall(pattern, weights) + if len(matches) != 1: + raise ValueError(f"Unexpected weights specification for the ViT-L backbone: {weights}") + hash = matches[0] + if hash == "eadcf0ff": + untie_global_and_local_cls_norm = True + kwargs["version"] = None + return _make_dinov3_vit( + img_size=224, + patch_size=16, + in_chans=3, + pos_embed_rope_base=100, + pos_embed_rope_normalize_coords="separate", + pos_embed_rope_rescale_coords=2, + pos_embed_rope_dtype="fp32", + embed_dim=1024, + depth=24, + num_heads=16, + ffn_ratio=4, + qkv_bias=True, + drop_path_rate=0.0, + layerscale_init=1.0e-05, + norm_layer="layernormbf16", + ffn_layer="mlp", + ffn_bias=True, + proj_bias=True, + n_storage_tokens=4, + mask_k_bias=True, + untie_global_and_local_cls_norm=untie_global_and_local_cls_norm, + pretrained=pretrained, + weights=weights, + compact_arch_name="vitl", + check_hash=check_hash, + **kwargs, + ) + + +def dinov3_vitl16plus( + *, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD1689M, + check_hash: bool = False, + **kwargs, +): + if "hash" not in kwargs: + kwargs["hash"] = "46503df0" + + return _make_dinov3_vit( + img_size=224, + patch_size=16, + in_chans=3, + pos_embed_rope_base=100, + pos_embed_rope_normalize_coords="separate", + pos_embed_rope_rescale_coords=2, + pos_embed_rope_dtype="fp32", + embed_dim=1024, + depth=24, + num_heads=16, + ffn_ratio=6.0, + qkv_bias=True, + drop_path_rate=0.0, + layerscale_init=1.0e-05, + norm_layer="layernormbf16", + ffn_layer="swiglu", + ffn_bias=True, + proj_bias=True, + n_storage_tokens=4, + mask_k_bias=True, + pretrained=pretrained, + weights=weights, + compact_arch_name="vitlplus", + check_hash=check_hash, + **kwargs, + ) + + +def dinov3_vith16plus( + *, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD1689M, + check_hash: bool = False, + **kwargs, +): + if "hash" not in kwargs: + kwargs["hash"] = "7c1da9a5" + + return _make_dinov3_vit( + img_size=224, + patch_size=16, + in_chans=3, + pos_embed_rope_base=100, + pos_embed_rope_normalize_coords="separate", + pos_embed_rope_rescale_coords=2, + pos_embed_rope_dtype="fp32", + embed_dim=1280, + depth=32, + num_heads=20, + ffn_ratio=6.0, + qkv_bias=True, + drop_path_rate=0.0, + layerscale_init=1.0e-05, + norm_layer="layernormbf16", + ffn_layer="swiglu", + ffn_bias=True, + proj_bias=True, + n_storage_tokens=4, + mask_k_bias=True, + pretrained=pretrained, + weights=weights, + compact_arch_name="vithplus", + check_hash=check_hash, + **kwargs, + ) + + +def dinov3_vit7b16( + *, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD1689M, + check_hash: bool = False, + **kwargs, +): + if weights == Weights.LVD1689M: + if "hash" not in kwargs: + kwargs["hash"] = "a955f4ea" + elif weights == Weights.SAT493M: + if "hash" not in kwargs: + kwargs["hash"] = "a6675841" + kwargs["version"] = None + untie_global_and_local_cls_norm = True + return _make_dinov3_vit( + img_size=224, + patch_size=16, + in_chans=3, + pos_embed_rope_base=100, + pos_embed_rope_normalize_coords="separate", + pos_embed_rope_rescale_coords=2, + pos_embed_rope_dtype="fp32", + embed_dim=4096, + depth=40, + num_heads=32, + ffn_ratio=3, + qkv_bias=False, + drop_path_rate=0.0, + layerscale_init=1.0e-05, + norm_layer="layernormbf16", + ffn_layer="swiglu64", + ffn_bias=True, + proj_bias=True, + n_storage_tokens=4, + mask_k_bias=True, + untie_global_and_local_cls_norm=untie_global_and_local_cls_norm, + pretrained=pretrained, + weights=weights, + compact_arch_name="vit7b", + check_hash=check_hash, + **kwargs, + ) + + +def dinov3_convnext_tiny( + *, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD1689M, + **kwargs, +): + _hash_convnext = "21b726bb" + if "hash" not in kwargs: + kwargs["hash"] = _hash_convnext + + from ..models.convnext import convnext_sizes + + size_dict = convnext_sizes["tiny"] + + model = _make_dinov3_convnext( + in_chans=3, + depths=size_dict["depths"], + dims=size_dict["dims"], + compact_arch_name="convnext_tiny", + drop_path_rate=0, + layer_scale_init_value=1e-6, + pretrained=pretrained, + weights=weights, + **kwargs, + ) + if not pretrained: + model.init_weights() + return model + + +def dinov3_convnext_small( + *, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD1689M, + **kwargs, +): + _hash_convnext = "296db49d" + if "hash" not in kwargs: + kwargs["hash"] = _hash_convnext + + from ..models.convnext import convnext_sizes + + size_dict = convnext_sizes["small"] + + model = _make_dinov3_convnext( + in_chans=3, + depths=size_dict["depths"], + dims=size_dict["dims"], + compact_arch_name="convnext_small", + drop_path_rate=0, + layer_scale_init_value=1e-6, + pretrained=pretrained, + weights=weights, + **kwargs, + ) + if not pretrained: + model.init_weights() + return model + + +def dinov3_convnext_base( + *, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD1689M, + **kwargs, +): + _hash_convnext = "801f2ba9" + if "hash" not in kwargs: + kwargs["hash"] = _hash_convnext + + from ..models.convnext import convnext_sizes + + size_dict = convnext_sizes["base"] + + model = _make_dinov3_convnext( + in_chans=3, + depths=size_dict["depths"], + dims=size_dict["dims"], + compact_arch_name="convnext_base", + drop_path_rate=0, + layer_scale_init_value=1e-6, + pretrained=pretrained, + weights=weights, + **kwargs, + ) + if not pretrained: + model.init_weights() + return model + + +def dinov3_convnext_large( + *, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD1689M, + **kwargs, +): + _hash_convnext = "61fa432d" + if "hash" not in kwargs: + kwargs["hash"] = _hash_convnext + + from ..models.convnext import convnext_sizes + + size_dict = convnext_sizes["large"] + + model = _make_dinov3_convnext( + in_chans=3, + depths=size_dict["depths"], + dims=size_dict["dims"], + compact_arch_name="convnext_large", + drop_path_rate=0, + layer_scale_init_value=1e-6, + pretrained=pretrained, + weights=weights, + **kwargs, + ) + if not pretrained: + model.init_weights() + return model diff --git a/kaloscope_dinov3/hub/classifiers.py b/kaloscope_dinov3/hub/classifiers.py new file mode 100644 index 0000000..ab53b3f --- /dev/null +++ b/kaloscope_dinov3/hub/classifiers.py @@ -0,0 +1,113 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +import os +from enum import Enum + +import torch +import torch.nn as nn + +from .backbones import ( + dinov3_vit7b16, + Weights as BackboneWeights, + convert_path_or_url_to_url, +) + +from .utils import _DINOV3_BASE_URL, _safe_load_state_dict_from_url + + +class ClassifierWeights(Enum): + IMAGENET1K = "IMAGENET1K" + + +def _make_dinov3_linear_classification_head( + *, + backbone_name: str = "dinov3_vit7b16", + embed_dim: int = 8192, + pretrained: bool = True, + classifier_weights: ClassifierWeights | str = ClassifierWeights.IMAGENET1K, + check_hash: bool = False, + **kwargs, +): + linear_head = nn.Linear(embed_dim, 1_000) + if pretrained: + if type(classifier_weights) is ClassifierWeights: + assert classifier_weights == ClassifierWeights.IMAGENET1K, ( + f"Unsupported weights for linear classifier: {classifier_weights}" + ) + weights_name = classifier_weights.value.lower() + hash = kwargs["hash"] if "hash" in kwargs else "90d8ed92" + model_filename = f"{backbone_name}_{weights_name}_linear_head-{hash}.pth" + url = os.path.join(_DINOV3_BASE_URL, backbone_name, model_filename) + else: + url = convert_path_or_url_to_url(classifier_weights) + state_dict = _safe_load_state_dict_from_url(url, map_location="cpu", check_hash=check_hash) + linear_head.load_state_dict(state_dict, strict=True) + return linear_head + + +class _LinearClassifierWrapper(nn.Module): + def __init__(self, *, backbone: nn.Module, linear_head: nn.Module): + super().__init__() + self.backbone = backbone + self.linear_head = linear_head + + def forward(self, x): + x = self.backbone.forward_features(x) + cls_token = x["x_norm_clstoken"] + patch_tokens = x["x_norm_patchtokens"] + linear_input = torch.cat( + [ + cls_token, + patch_tokens.mean(dim=1), + ], + dim=1, + ) + return self.linear_head(linear_input) + + +def _make_dinov3_linear_classifier( + *, + backbone_name: str = "dinov3_vit7b16", + pretrained: bool = True, + classifier_weights: ClassifierWeights | str = ClassifierWeights.IMAGENET1K, + backbone_weights: BackboneWeights | str = BackboneWeights.LVD1689M, + check_hash: bool = False, + **kwargs, +): + if backbone_name == "dinov3_vit7b16": + backbone = dinov3_vit7b16(pretrained=pretrained, weights=backbone_weights, check_hash=check_hash) + else: + raise AssertionError(f"Unsupported backbone: {backbone_name}, linear classifiers are provided only for ViT-7b") + embed_dim = backbone.embed_dim + linear_head = _make_dinov3_linear_classification_head( + backbone_name=backbone_name, + embed_dim=2 * embed_dim, + pretrained=pretrained, + classifier_weights=classifier_weights, + **kwargs, + ) + return _LinearClassifierWrapper(backbone=backbone, linear_head=linear_head) + + +def dinov3_vit7b16_lc( + *, + pretrained: bool = True, + weights: ClassifierWeights | str = ClassifierWeights.IMAGENET1K, + backbone_weights: BackboneWeights | str = BackboneWeights.LVD1689M, + check_hash: bool = False, + **kwargs, +): + """ + Linear classifier on top of a DINOv3 ViT-7B/16 backbone pretrained on the LVD-1689M dataset and trained on ImageNet-1k. + """ + return _make_dinov3_linear_classifier( + backbone_name="dinov3_vit7b16", + pretrained=pretrained, + classifier_weights=weights, + backbone_weights=backbone_weights, + check_hash=check_hash, + **kwargs, + ) diff --git a/kaloscope_dinov3/hub/utils.py b/kaloscope_dinov3/hub/utils.py new file mode 100644 index 0000000..5c82f32 --- /dev/null +++ b/kaloscope_dinov3/hub/utils.py @@ -0,0 +1,18 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +import torch + + +_DINOV3_BASE_URL = "https://dl.fbaipublicfiles.com/dinov3" + + +def _safe_load_state_dict_from_url(url: str, **kwargs): + # See https://github.com/pytorch/pytorch/releases/tag/v2.1.0 (Misc / #98479) + if torch.__version__ >= (2, 1): + local_kwargs = {**kwargs, "weights_only": True} + else: + local_kwargs = kwargs + return torch.hub.load_state_dict_from_url(url, **local_kwargs) diff --git a/kaloscope_dinov3/layers/__init__.py b/kaloscope_dinov3/layers/__init__.py new file mode 100644 index 0000000..68fdc17 --- /dev/null +++ b/kaloscope_dinov3/layers/__init__.py @@ -0,0 +1,12 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +from .attention import CausalSelfAttention, LinearKMaskedBias, SelfAttention +from .block import CausalSelfAttentionBlock, SelfAttentionBlock +from .ffn_layers import Mlp, SwiGLUFFN +from .layer_scale import LayerScale +from .patch_embed import PatchEmbed +from .rms_norm import RMSNorm +from .rope_position_encoding import RopePositionEmbedding diff --git a/kaloscope_dinov3/layers/attention.py b/kaloscope_dinov3/layers/attention.py new file mode 100644 index 0000000..ac388e7 --- /dev/null +++ b/kaloscope_dinov3/layers/attention.py @@ -0,0 +1,164 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +import math +from typing import List, Tuple + +import torch +import torch.nn.functional as F +from kaloscope_dinov3.utils import cat_keep_shapes, uncat_with_shapes +from torch import Tensor, nn + + +# RoPE-related functions: +def rope_rotate_half(x: Tensor) -> Tensor: + # x: [ x0 x1 x2 x3 x4 x5] + # out: [-x3 -x4 -x5 x0 x1 x2] + x1, x2 = x.chunk(2, dim=-1) + return torch.cat([-x2, x1], dim=-1) + + +def rope_apply(x: Tensor, sin: Tensor, cos: Tensor) -> Tensor: + # x: [..., D], eg [x0, x1, x2, x3, x4, x5] + # sin: [..., D], eg [sin0, sin1, sin2, sin0, sin1, sin2] + # cos: [..., D], eg [cos0, cos1, cos2, cos0, cos1, cos2] + return (x * cos) + (rope_rotate_half(x) * sin) + + +class LinearKMaskedBias(nn.Linear): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + o = self.out_features + assert o % 3 == 0 + if self.bias is not None: + self.register_buffer("bias_mask", torch.full_like(self.bias, fill_value=math.nan)) + + def forward(self, input: Tensor) -> Tensor: + masked_bias = self.bias * self.bias_mask.to(self.bias.dtype) if self.bias is not None else None + return F.linear(input, self.weight, masked_bias) + + +class SelfAttention(nn.Module): + def __init__( + self, + dim: int, + num_heads: int = 8, + qkv_bias: bool = False, + proj_bias: bool = True, + attn_drop: float = 0.0, + proj_drop: float = 0.0, + mask_k_bias: bool = False, + device=None, + ) -> None: + super().__init__() + self.num_heads = num_heads + head_dim = dim // num_heads + self.scale = head_dim**-0.5 + + linear_class = LinearKMaskedBias if mask_k_bias else nn.Linear + self.qkv = linear_class(dim, dim * 3, bias=qkv_bias, device=device) + self.attn_drop = nn.Dropout(attn_drop) + self.proj = nn.Linear(dim, dim, bias=proj_bias, device=device) + self.proj_drop = nn.Dropout(proj_drop) + + def apply_rope(self, q: Tensor, k: Tensor, rope: Tensor | Tuple[Tensor, Tensor]) -> Tuple[Tensor, Tensor]: + # All operations will use the dtype of rope, the output is cast back to the dtype of q and k + q_dtype = q.dtype + k_dtype = k.dtype + sin, cos = rope + rope_dtype = sin.dtype + q = q.to(dtype=rope_dtype) + k = k.to(dtype=rope_dtype) + N = q.shape[-2] + prefix = N - sin.shape[-2] + assert prefix >= 0 + q_prefix = q[:, :, :prefix, :] + q = rope_apply(q[:, :, prefix:, :], sin, cos) # [B, head, hw, D//head] + q = torch.cat((q_prefix, q), dim=-2) # [B, head, N, D//head] + k_prefix = k[:, :, :prefix, :] + k = rope_apply(k[:, :, prefix:, :], sin, cos) # [B, head, hw, D//head] + k = torch.cat((k_prefix, k), dim=-2) # [B, head, N, D//head] + q = q.to(dtype=q_dtype) + k = k.to(dtype=k_dtype) + return q, k + + def forward(self, x: Tensor, attn_bias=None, rope: Tensor = None) -> Tensor: + qkv = self.qkv(x) + attn_v = self.compute_attention(qkv=qkv, attn_bias=attn_bias, rope=rope) + x = self.proj(attn_v) + x = self.proj_drop(x) + return x + + def forward_list(self, x_list, attn_bias=None, rope_list=None) -> List[Tensor]: + assert len(x_list) == len(rope_list) # should be enforced by the Block + x_flat, shapes, num_tokens = cat_keep_shapes(x_list) + qkv_flat = self.qkv(x_flat) + qkv_list = uncat_with_shapes(qkv_flat, shapes, num_tokens) + att_out = [] + for _, (qkv, _, rope) in enumerate(zip(qkv_list, shapes, rope_list)): + att_out.append(self.compute_attention(qkv, attn_bias=attn_bias, rope=rope)) + x_flat, shapes, num_tokens = cat_keep_shapes(att_out) + x_flat = self.proj(x_flat) + return uncat_with_shapes(x_flat, shapes, num_tokens) + + def compute_attention(self, qkv: Tensor, attn_bias=None, rope=None) -> Tensor: + assert attn_bias is None + B, N, _ = qkv.shape + C = self.qkv.in_features + + qkv = qkv.reshape(B, N, 3, self.num_heads, C // self.num_heads) + q, k, v = torch.unbind(qkv, 2) + q, k, v = [t.transpose(1, 2) for t in [q, k, v]] + if rope is not None: + q, k = self.apply_rope(q, k, rope) + x = torch.nn.functional.scaled_dot_product_attention(q, k, v) + x = x.transpose(1, 2) + return x.reshape([B, N, C]) + + +class CausalSelfAttention(nn.Module): + def __init__( + self, + dim: int, + num_heads: int = 8, + qkv_bias: bool = False, + proj_bias: bool = True, + attn_drop: float = 0.0, + proj_drop: float = 0.0, + ) -> None: + super().__init__() + self.dim = dim + self.num_heads = num_heads + head_dim = dim // num_heads + self.scale = head_dim**-0.5 + + self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) + self.attn_drop = attn_drop + self.proj = nn.Linear(dim, dim, bias=proj_bias) + self.proj_drop = nn.Dropout(proj_drop) + + def init_weights( + self, init_attn_std: float | None = None, init_proj_std: float | None = None, factor: float = 1.0 + ) -> None: + init_attn_std = init_attn_std or (self.dim**-0.5) + init_proj_std = init_proj_std or init_attn_std * factor + nn.init.normal_(self.qkv.weight, std=init_attn_std) + nn.init.normal_(self.proj.weight, std=init_proj_std) + if self.qkv.bias is not None: + nn.init.zeros_(self.qkv.bias) + if self.proj.bias is not None: + nn.init.zeros_(self.proj.bias) + + def forward(self, x: Tensor, is_causal: bool = True) -> Tensor: + B, N, C = x.shape + qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads) + q, k, v = torch.unbind(qkv, 2) + q, k, v = [t.transpose(1, 2) for t in [q, k, v]] + x = torch.nn.functional.scaled_dot_product_attention( + q, k, v, attn_mask=None, dropout_p=self.attn_drop if self.training else 0, is_causal=is_causal + ) + x = x.transpose(1, 2).contiguous().view(B, N, C) + x = self.proj_drop(self.proj(x)) + return x diff --git a/kaloscope_dinov3/layers/block.py b/kaloscope_dinov3/layers/block.py new file mode 100644 index 0000000..f27d4f1 --- /dev/null +++ b/kaloscope_dinov3/layers/block.py @@ -0,0 +1,269 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +from typing import Callable, List, Optional + +import torch +from torch import Tensor, nn + +from kaloscope_dinov3.utils import cat_keep_shapes, uncat_with_shapes + +from .attention import CausalSelfAttention, SelfAttention +from .ffn_layers import Mlp +from .layer_scale import LayerScale # , DropPath + +torch._dynamo.config.automatic_dynamic_shapes = False +torch._dynamo.config.accumulated_cache_size_limit = 1024 + + +class SelfAttentionBlock(nn.Module): + def __init__( + self, + dim: int, + num_heads: int, + ffn_ratio: float = 4.0, + qkv_bias: bool = False, + proj_bias: bool = True, + ffn_bias: bool = True, + drop: float = 0.0, + attn_drop: float = 0.0, + init_values=None, + drop_path: float = 0.0, + act_layer: Callable[..., nn.Module] = nn.GELU, + norm_layer: Callable[..., nn.Module] = nn.LayerNorm, + attn_class: Callable[..., nn.Module] = SelfAttention, + ffn_layer: Callable[..., nn.Module] = Mlp, + mask_k_bias: bool = False, + device=None, + ) -> None: + super().__init__() + # print(f"biases: qkv: {qkv_bias}, proj: {proj_bias}, ffn: {ffn_bias}") + self.norm1 = norm_layer(dim) + self.attn = attn_class( + dim, + num_heads=num_heads, + qkv_bias=qkv_bias, + proj_bias=proj_bias, + attn_drop=attn_drop, + proj_drop=drop, + mask_k_bias=mask_k_bias, + device=device, + ) + self.ls1 = LayerScale(dim, init_values=init_values, device=device) if init_values else nn.Identity() + + self.norm2 = norm_layer(dim) + mlp_hidden_dim = int(dim * ffn_ratio) + self.mlp = ffn_layer( + in_features=dim, + hidden_features=mlp_hidden_dim, + act_layer=act_layer, + drop=drop, + bias=ffn_bias, + device=device, + ) + self.ls2 = LayerScale(dim, init_values=init_values, device=device) if init_values else nn.Identity() + + self.sample_drop_ratio = drop_path + + @staticmethod + def _maybe_index_rope(rope: tuple[Tensor, Tensor] | None, indices: Tensor) -> tuple[Tensor, Tensor] | None: + if rope is None: + return None + + sin, cos = rope + assert sin.ndim == cos.ndim + if sin.ndim == 4: + # If the rope embedding has a batch dimension (is different for each batch element), index into it + return sin[indices], cos[indices] # [batch, heads, patches, embed_dim] + else: + # No batch dimension, do not index + return sin, cos # [heads, patches, embed_dim] or [patches, embed_dim] + + def _forward(self, x: Tensor, rope=None) -> Tensor: + """ + This is the reference implementation for a single tensor, matching what is done below for a list. + We call the list op on [x] instead of this function. + """ + b, _, _ = x.shape + sample_subset_size = max(int(b * (1 - self.sample_drop_ratio)), 1) + residual_scale_factor = b / sample_subset_size + + if self.training and self.sample_drop_ratio > 0.0: + indices_1 = (torch.randperm(b, device=x.device))[:sample_subset_size] + + x_subset_1 = x[indices_1] + rope_subset = self._maybe_index_rope(rope, indices_1) + residual_1 = self.attn(self.norm1(x_subset_1), rope=rope_subset) + + x_attn = torch.index_add( + x, + dim=0, + source=self.ls1(residual_1), + index=indices_1, + alpha=residual_scale_factor, + ) + + indices_2 = (torch.randperm(b, device=x.device))[:sample_subset_size] + + x_subset_2 = x_attn[indices_2] + residual_2 = self.mlp(self.norm2(x_subset_2)) + + x_ffn = torch.index_add( + x_attn, + dim=0, + source=self.ls2(residual_2), + index=indices_2, + alpha=residual_scale_factor, + ) + else: + x_attn = x + self.ls1(self.attn(self.norm1(x), rope=rope)) + x_ffn = x_attn + self.ls2(self.mlp(self.norm2(x_attn))) + + return x_ffn + + def _forward_list(self, x_list: List[Tensor], rope_list=None) -> List[Tensor]: + """ + This list operator concatenates the tokens from the list of inputs together to save + on the elementwise operations. Torch-compile memory-planning allows hiding the overhead + related to concat ops. + """ + b_list = [x.shape[0] for x in x_list] + sample_subset_sizes = [max(int(b * (1 - self.sample_drop_ratio)), 1) for b in b_list] + residual_scale_factors = [b / sample_subset_size for b, sample_subset_size in zip(b_list, sample_subset_sizes)] + + if self.training and self.sample_drop_ratio > 0.0: + indices_1_list = [ + (torch.randperm(b, device=x.device))[:sample_subset_size] + for x, b, sample_subset_size in zip(x_list, b_list, sample_subset_sizes) + ] + x_subset_1_list = [x[indices_1] for x, indices_1 in zip(x_list, indices_1_list)] + + if rope_list is not None: + rope_subset_list = [ + self._maybe_index_rope(rope, indices_1) for rope, indices_1 in zip(rope_list, indices_1_list) + ] + else: + rope_subset_list = rope_list + + flattened, shapes, num_tokens = cat_keep_shapes(x_subset_1_list) + norm1 = uncat_with_shapes(self.norm1(flattened), shapes, num_tokens) + residual_1_list = self.attn.forward_list(norm1, rope_list=rope_subset_list) + + x_attn_list = [ + torch.index_add( + x, + dim=0, + source=self.ls1(residual_1), + index=indices_1, + alpha=residual_scale_factor, + ) + for x, residual_1, indices_1, residual_scale_factor in zip( + x_list, residual_1_list, indices_1_list, residual_scale_factors + ) + ] + + indices_2_list = [ + (torch.randperm(b, device=x.device))[:sample_subset_size] + for x, b, sample_subset_size in zip(x_list, b_list, sample_subset_sizes) + ] + x_subset_2_list = [x[indices_2] for x, indices_2 in zip(x_attn_list, indices_2_list)] + flattened, shapes, num_tokens = cat_keep_shapes(x_subset_2_list) + norm2_flat = self.norm2(flattened) + norm2_list = uncat_with_shapes(norm2_flat, shapes, num_tokens) + + residual_2_list = self.mlp.forward_list(norm2_list) + + x_ffn = [ + torch.index_add( + x_attn, + dim=0, + source=self.ls2(residual_2), + index=indices_2, + alpha=residual_scale_factor, + ) + for x_attn, residual_2, indices_2, residual_scale_factor in zip( + x_attn_list, residual_2_list, indices_2_list, residual_scale_factors + ) + ] + else: + x_out = [] + for x, rope in zip(x_list, rope_list): + x_attn = x + self.ls1(self.attn(self.norm1(x), rope=rope)) + x_ffn = x_attn + self.ls2(self.mlp(self.norm2(x_attn))) + x_out.append(x_ffn) + x_ffn = x_out + + return x_ffn + + def forward(self, x_or_x_list, rope_or_rope_list=None) -> List[Tensor]: + if isinstance(x_or_x_list, Tensor): + # for reference: + # return self._forward(x_or_x_list, rope=rope_or_rope_list) + # in order to match implementations we call the list op: + return self._forward_list([x_or_x_list], rope_list=[rope_or_rope_list])[0] + elif isinstance(x_or_x_list, list): + if rope_or_rope_list is None: + rope_or_rope_list = [None for x in x_or_x_list] + # return [self._forward(x, rope=rope) for x, rope in zip(x_or_x_list, rope_or_rope_list)] + return self._forward_list(x_or_x_list, rope_list=rope_or_rope_list) + else: + raise AssertionError + + +class CausalSelfAttentionBlock(nn.Module): + def __init__( + self, + dim: int, + num_heads: int, + ffn_ratio: float = 4.0, + ls_init_value: Optional[float] = None, + is_causal: bool = True, + act_layer: Callable = nn.GELU, + norm_layer: Callable = nn.LayerNorm, + dropout_prob: float = 0.0, + ): + super().__init__() + + self.dim = dim + self.is_causal = is_causal + self.ls1 = LayerScale(dim, init_values=ls_init_value) if ls_init_value else nn.Identity() + self.attention_norm = norm_layer(dim) + self.attention = CausalSelfAttention(dim, num_heads, attn_drop=dropout_prob, proj_drop=dropout_prob) + + self.ffn_norm = norm_layer(dim) + ffn_hidden_dim = int(dim * ffn_ratio) + self.feed_forward = Mlp( + in_features=dim, + hidden_features=ffn_hidden_dim, + drop=dropout_prob, + act_layer=act_layer, + ) + + self.ls2 = LayerScale(dim, init_values=ls_init_value) if ls_init_value else nn.Identity() + + def init_weights( + self, + init_attn_std: float | None = None, + init_proj_std: float | None = None, + init_fc_std: float | None = None, + factor: float = 1.0, + ) -> None: + init_attn_std = init_attn_std or (self.dim**-0.5) + init_proj_std = init_proj_std or init_attn_std * factor + init_fc_std = init_fc_std or (2 * self.dim) ** -0.5 + self.attention.init_weights(init_attn_std, init_proj_std) + self.attention_norm.reset_parameters() + nn.init.normal_(self.feed_forward.fc1.weight, std=init_fc_std) + nn.init.normal_(self.feed_forward.fc2.weight, std=init_proj_std) + self.ffn_norm.reset_parameters() + + def forward( + self, + x: torch.Tensor, + ): + + x_attn = x + self.ls1(self.attention(self.attention_norm(x), self.is_causal)) + x_ffn = x_attn + self.ls2(self.feed_forward(self.ffn_norm(x_attn))) + return x_ffn diff --git a/kaloscope_dinov3/layers/ffn_layers.py b/kaloscope_dinov3/layers/ffn_layers.py new file mode 100644 index 0000000..1c25e68 --- /dev/null +++ b/kaloscope_dinov3/layers/ffn_layers.py @@ -0,0 +1,77 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +from typing import Callable, List, Optional + +import torch.nn.functional as F +from torch import Tensor, nn + +from kaloscope_dinov3.utils import cat_keep_shapes, uncat_with_shapes + + +class ListForwardMixin(object): + def forward(self, x: Tensor): + raise NotImplementedError + + def forward_list(self, x_list: List[Tensor]) -> List[Tensor]: + x_flat, shapes, num_tokens = cat_keep_shapes(x_list) + x_flat = self.forward(x_flat) + return uncat_with_shapes(x_flat, shapes, num_tokens) + + +class Mlp(nn.Module, ListForwardMixin): + def __init__( + self, + in_features: int, + hidden_features: Optional[int] = None, + out_features: Optional[int] = None, + act_layer: Callable[..., nn.Module] = nn.GELU, + drop: float = 0.0, + bias: bool = True, + device=None, + ) -> None: + super().__init__() + out_features = out_features or in_features + hidden_features = hidden_features or in_features + self.fc1 = nn.Linear(in_features, hidden_features, bias=bias, device=device) + self.act = act_layer() + self.fc2 = nn.Linear(hidden_features, out_features, bias=bias, device=device) + self.drop = nn.Dropout(drop) + + def forward(self, x: Tensor) -> Tensor: + x = self.fc1(x) + x = self.act(x) + x = self.drop(x) + x = self.fc2(x) + x = self.drop(x) + return x + + +class SwiGLUFFN(nn.Module, ListForwardMixin): + def __init__( + self, + in_features: int, + hidden_features: Optional[int] = None, + out_features: Optional[int] = None, + act_layer: Optional[Callable[..., nn.Module]] = None, + drop: float = 0.0, + bias: bool = True, + align_to: int = 8, + device=None, + ) -> None: + super().__init__() + out_features = out_features or in_features + hidden_features = hidden_features or in_features + d = int(hidden_features * 2 / 3) + swiglu_hidden_features = d + (-d % align_to) + self.w1 = nn.Linear(in_features, swiglu_hidden_features, bias=bias, device=device) + self.w2 = nn.Linear(in_features, swiglu_hidden_features, bias=bias, device=device) + self.w3 = nn.Linear(swiglu_hidden_features, out_features, bias=bias, device=device) + + def forward(self, x: Tensor) -> Tensor: + x1 = self.w1(x) + x2 = self.w2(x) + hidden = F.silu(x1) * x2 + return self.w3(hidden) diff --git a/kaloscope_dinov3/layers/layer_scale.py b/kaloscope_dinov3/layers/layer_scale.py new file mode 100644 index 0000000..0b72b7c --- /dev/null +++ b/kaloscope_dinov3/layers/layer_scale.py @@ -0,0 +1,29 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +from typing import Union + +import torch +from torch import Tensor, nn + + +class LayerScale(nn.Module): + def __init__( + self, + dim: int, + init_values: Union[float, Tensor] = 1e-5, + inplace: bool = False, + device=None, + ) -> None: + super().__init__() + self.inplace = inplace + self.gamma = nn.Parameter(torch.empty(dim, device=device)) + self.init_values = init_values + + def reset_parameters(self): + nn.init.constant_(self.gamma, self.init_values) + + def forward(self, x: Tensor) -> Tensor: + return x.mul_(self.gamma) if self.inplace else x * self.gamma diff --git a/kaloscope_dinov3/layers/patch_embed.py b/kaloscope_dinov3/layers/patch_embed.py new file mode 100644 index 0000000..760343f --- /dev/null +++ b/kaloscope_dinov3/layers/patch_embed.py @@ -0,0 +1,89 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +import math +from typing import Callable, Tuple, Union + +from torch import Tensor, nn + + +def make_2tuple(x): + if isinstance(x, tuple): + assert len(x) == 2 + return x + + assert isinstance(x, int) + return (x, x) + + +class PatchEmbed(nn.Module): + """ + 2D image to patch embedding: (B,C,H,W) -> (B,N,D) + + Args: + img_size: Image size. + patch_size: Patch token size. + in_chans: Number of input image channels. + embed_dim: Number of linear projection output channels. + norm_layer: Normalization layer. + """ + + def __init__( + self, + img_size: Union[int, Tuple[int, int]] = 224, + patch_size: Union[int, Tuple[int, int]] = 16, + in_chans: int = 3, + embed_dim: int = 768, + norm_layer: Callable | None = None, + flatten_embedding: bool = True, + ) -> None: + super().__init__() + + image_HW = make_2tuple(img_size) + patch_HW = make_2tuple(patch_size) + patch_grid_size = ( + image_HW[0] // patch_HW[0], + image_HW[1] // patch_HW[1], + ) + + self.img_size = image_HW + self.patch_size = patch_HW + self.patches_resolution = patch_grid_size + self.num_patches = patch_grid_size[0] * patch_grid_size[1] + + self.in_chans = in_chans + self.embed_dim = embed_dim + + self.flatten_embedding = flatten_embedding + + self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_HW, stride=patch_HW) + self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity() + + def forward(self, x: Tensor) -> Tensor: + _, _, H, W = x.shape + # patch_H, patch_W = self.patch_size + # assert H % patch_H == 0, f"Input image height {H} is not a multiple of patch height {patch_H}" + # assert W % patch_W == 0, f"Input image width {W} is not a multiple of patch width: {patch_W}" + + x = self.proj(x) # B C H W + H, W = x.size(2), x.size(3) + x = x.flatten(2).transpose(1, 2) # B HW C + x = self.norm(x) + if not self.flatten_embedding: + x = x.reshape(-1, H, W, self.embed_dim) # B H W C + return x + + def flops(self) -> float: + Ho, Wo = self.patches_resolution + flops = Ho * Wo * self.embed_dim * self.in_chans * (self.patch_size[0] * self.patch_size[1]) + if self.norm is not None: + flops += Ho * Wo * self.embed_dim + return flops + + def reset_parameters(self): + k = 1 / (self.in_chans * (self.patch_size[0] ** 2)) + nn.init.uniform_(self.proj.weight, -math.sqrt(k), math.sqrt(k)) + if self.proj.bias is not None: + nn.init.uniform_(self.proj.bias, -math.sqrt(k), math.sqrt(k)) diff --git a/kaloscope_dinov3/layers/rms_norm.py b/kaloscope_dinov3/layers/rms_norm.py new file mode 100644 index 0000000..1d0a89c --- /dev/null +++ b/kaloscope_dinov3/layers/rms_norm.py @@ -0,0 +1,24 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +import torch +from torch import Tensor, nn + + +class RMSNorm(nn.Module): + def __init__(self, dim: int, eps: float = 1e-5): + super().__init__() + self.weight = nn.Parameter(torch.ones(dim)) + self.eps = eps + + def reset_parameters(self) -> None: + nn.init.constant_(self.weight, 1) + + def _norm(self, x: Tensor) -> Tensor: + return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) + + def forward(self, x: Tensor) -> Tensor: + output = self._norm(x.float()).type_as(x) + return output * self.weight diff --git a/kaloscope_dinov3/layers/rope_position_encoding.py b/kaloscope_dinov3/layers/rope_position_encoding.py new file mode 100644 index 0000000..2635d09 --- /dev/null +++ b/kaloscope_dinov3/layers/rope_position_encoding.py @@ -0,0 +1,121 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +import math +from typing import Literal + +import numpy as np +import torch +from torch import Tensor, nn + + +# RoPE positional embedding with no mixing of coordinates (axial) and no learnable weights +# Supports two parametrizations of the rope parameters: either using `base` or `min_period` and `max_period`. +class RopePositionEmbedding(nn.Module): + def __init__( + self, + embed_dim: int, + *, + num_heads: int, + base: float | None = 100.0, + min_period: float | None = None, + max_period: float | None = None, + normalize_coords: Literal["min", "max", "separate"] = "separate", + shift_coords: float | None = None, + jitter_coords: float | None = None, + rescale_coords: float | None = None, + dtype: torch.dtype | None = None, + device: torch.device | None = None, + ): + super().__init__() + assert embed_dim % (4 * num_heads) == 0 + both_periods = min_period is not None and max_period is not None + if (base is None and not both_periods) or (base is not None and both_periods): + raise ValueError("Either `base` or `min_period`+`max_period` must be provided.") + + D_head = embed_dim // num_heads + self.base = base + self.min_period = min_period + self.max_period = max_period + self.D_head = D_head + self.normalize_coords = normalize_coords + self.shift_coords = shift_coords + self.jitter_coords = jitter_coords + self.rescale_coords = rescale_coords + + # Needs persistent=True because we do teacher.load_state_dict(student.state_dict()) to initialize the teacher + self.dtype = dtype # Don't rely on self.periods.dtype + self.register_buffer( + "periods", + torch.empty(D_head // 4, device=device, dtype=dtype), + persistent=True, + ) + self._init_weights() + + def forward(self, *, H: int, W: int) -> tuple[Tensor, Tensor]: + device = self.periods.device + dtype = self.dtype + dd = {"device": device, "dtype": dtype} + + # Prepare coords in range [-1, +1] + if self.normalize_coords == "max": + max_HW = max(H, W) + coords_h = torch.arange(0.5, H, **dd) / max_HW # [H] + coords_w = torch.arange(0.5, W, **dd) / max_HW # [W] + elif self.normalize_coords == "min": + min_HW = min(H, W) + coords_h = torch.arange(0.5, H, **dd) / min_HW # [H] + coords_w = torch.arange(0.5, W, **dd) / min_HW # [W] + elif self.normalize_coords == "separate": + coords_h = torch.arange(0.5, H, **dd) / H # [H] + coords_w = torch.arange(0.5, W, **dd) / W # [W] + else: + raise ValueError(f"Unknown normalize_coords: {self.normalize_coords}") + coords = torch.stack(torch.meshgrid(coords_h, coords_w, indexing="ij"), dim=-1) # [H, W, 2] + coords = coords.flatten(0, 1) # [HW, 2] + coords = 2.0 * coords - 1.0 # Shift range [0, 1] to [-1, +1] + + # Shift coords by adding a uniform value in [-shift, shift] + if self.training and self.shift_coords is not None: + shift_hw = torch.empty(2, **dd).uniform_(-self.shift_coords, self.shift_coords) + coords += shift_hw[None, :] + + # Jitter coords by multiplying the range [-1, 1] by a log-uniform value in [1/jitter, jitter] + if self.training and self.jitter_coords is not None: + jitter_max = np.log(self.jitter_coords) + jitter_min = -jitter_max + jitter_hw = torch.empty(2, **dd).uniform_(jitter_min, jitter_max).exp() + coords *= jitter_hw[None, :] + + # Rescale coords by multiplying the range [-1, 1] by a log-uniform value in [1/rescale, rescale] + if self.training and self.rescale_coords is not None: + rescale_max = np.log(self.rescale_coords) + rescale_min = -rescale_max + rescale_hw = torch.empty(1, **dd).uniform_(rescale_min, rescale_max).exp() + coords *= rescale_hw + + # Prepare angles and sin/cos + angles = 2 * math.pi * coords[:, :, None] / self.periods[None, None, :] # [HW, 2, D//4] + angles = angles.flatten(1, 2) # [HW, D//2] + angles = angles.tile(2) # [HW, D] + cos = torch.cos(angles) # [HW, D] + sin = torch.sin(angles) # [HW, D] + + return (sin, cos) # 2 * [HW, D] + + def _init_weights(self): + device = self.periods.device + dtype = self.dtype + if self.base is not None: + periods = self.base ** ( + 2 * torch.arange(self.D_head // 4, device=device, dtype=dtype) / (self.D_head // 2) + ) # [D//4] + else: + base = self.max_period / self.min_period + exponents = torch.linspace(0, 1, self.D_head // 4, device=device, dtype=dtype) # [D//4] range [0, 1] + periods = base**exponents # range [1, max_period / min_period] + periods = periods / base # range [min_period / max_period, 1] + periods = periods * self.max_period # range [min_period, max_period] + self.periods.data = periods diff --git a/kaloscope_dinov3/preprocessing.py b/kaloscope_dinov3/preprocessing.py new file mode 100644 index 0000000..c641bc0 --- /dev/null +++ b/kaloscope_dinov3/preprocessing.py @@ -0,0 +1,67 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +"""Deterministic image preprocessing for inference; no datasets or augmentations.""" +import importlib + +import torch +from PIL import Image, ImageOps +from torchvision import transforms + + +def prepare_rgb(image): + """Apply orientation and composite transparency before discarding alpha.""" + image = ImageOps.exif_transpose(image) + if "A" in image.getbands() or "transparency" in image.info: + rgba = image.convert("RGBA") + background = Image.new("RGBA", rgba.size, (255, 255, 255, 255)) + background.alpha_composite(rgba) + return background.convert("RGB") + return image.convert("RGB") + + +def _import_callable(dotted_path): + # Existing checkpoints can refer to the original preprocessing module. + for prefix in ("dinov3.finetune.data.", "kaloscope_dinov3.finetune.data."): + if dotted_path.startswith(prefix): + dotted_path = "kaloscope_dinov3.preprocessing." + dotted_path[len(prefix):] + break + module_path, _, attr_name = dotted_path.rpartition(".") + if not module_path: + raise ValueError(f"Invalid transform path: {dotted_path!r} (must be 'module.attr')") + return getattr(importlib.import_module(module_path), attr_name) + + +def image_transform(size, mean=None, std=None, custom_transform=None): + if custom_transform is not None: + factory = _import_callable(custom_transform) + return transforms.Compose([prepare_rgb, factory(resize_size=size)]) + return transforms.Compose([ + prepare_rgb, + transforms.Resize(round(size * 256 / 224)), + transforms.CenterCrop(size), + transforms.ToTensor(), + transforms.Normalize(mean or [0.485, 0.456, 0.406], std or [0.229, 0.224, 0.225]), + ]) + + +def lvd_transform(resize_size=256): + from torchvision.transforms import v2 + return v2.Compose([ + v2.ToImage(), + v2.Resize((resize_size, resize_size), antialias=True), + v2.ToDtype(torch.float32, scale=True), + v2.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)), + ]) + + +def sat_transform(resize_size=256): + from torchvision.transforms import v2 + return v2.Compose([ + v2.ToImage(), + v2.Resize((resize_size, resize_size), antialias=True), + v2.ToDtype(torch.float32, scale=True), + v2.Normalize(mean=(0.430, 0.411, 0.296), std=(0.213, 0.156, 0.143)), + ]) diff --git a/kaloscope_dinov3/utils/__init__.py b/kaloscope_dinov3/utils/__init__.py new file mode 100644 index 0000000..245f63e --- /dev/null +++ b/kaloscope_dinov3/utils/__init__.py @@ -0,0 +1,10 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +from .utils import ( + cat_keep_shapes, + named_apply, + uncat_with_shapes, +) diff --git a/kaloscope_dinov3/utils/utils.py b/kaloscope_dinov3/utils/utils.py new file mode 100644 index 0000000..556369a --- /dev/null +++ b/kaloscope_dinov3/utils/utils.py @@ -0,0 +1,46 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This software may be used and distributed in accordance with +# the terms of the DINOv3 License Agreement. + +from typing import Callable, List, Tuple + +import torch +from torch import Tensor, nn + + +def cat_keep_shapes(x_list: List[Tensor]) -> Tuple[Tensor, List[Tuple[int]], List[int]]: + shapes = [x.shape for x in x_list] + num_tokens = [x.select(dim=-1, index=0).numel() for x in x_list] + flattened = torch.cat([x.flatten(0, -2) for x in x_list]) + return flattened, shapes, num_tokens + + +def uncat_with_shapes(flattened: Tensor, shapes: List[Tuple[int]], num_tokens: List[int]) -> List[Tensor]: + outputs_splitted = torch.split_with_sizes(flattened, num_tokens, dim=0) + shapes_adjusted = [shape[:-1] + torch.Size([flattened.shape[-1]]) for shape in shapes] + outputs_reshaped = [o.reshape(shape) for o, shape in zip(outputs_splitted, shapes_adjusted)] + return outputs_reshaped + + +def named_apply( + fn: Callable, + module: nn.Module, + name: str = "", + depth_first: bool = True, + include_root: bool = False, +) -> nn.Module: + if not depth_first and include_root: + fn(module=module, name=name) + for child_name, child_module in module.named_children(): + child_name = ".".join((name, child_name)) if name else child_name + named_apply( + fn=fn, + module=child_module, + name=child_name, + depth_first=depth_first, + include_root=True, + ) + if depth_first and include_root: + fn(module=module, name=name) + return module diff --git a/model_loading.py b/model_loading.py new file mode 100644 index 0000000..c8e94c1 --- /dev/null +++ b/model_loading.py @@ -0,0 +1,339 @@ +"""Shared offline model loading for ComfyUI, CLI, WebUI and API.""" +import csv +import json +from pathlib import Path + +import torch +from torch import nn + +LSNET_MODELS = ( + "lsnet_t_artist", "lsnet_s_artist", "lsnet_b_artist", "lsnet_l_artist", + "lsnet_xl_artist", "lsnet_xl_artist_448", +) +CHECKPOINT_EXTENSIONS = (".pt", ".pth", ".ckpt", ".safetensors") +FEATURE_OUTPUTS = ( + "default", "backbone", "cls", "mean", "cls_mean", "projector", + "patch_tokens", "patch_map", "storage_tokens", "all_tokens", "prenorm", + "intermediate_cls", "intermediate_mean", "intermediate_cls_mean", + "intermediate_patch_tokens", "intermediate_patch_map", "intermediate_storage_tokens", + "intermediate_all_tokens", "intermediate_prenorm", +) + + +def _selected_layers(text, count): + try: + requested = [int(item.strip()) for item in text.split(",")] + except (ValueError, AttributeError): + raise ValueError("layers must be comma-separated indices, e.g. -1 or 8,9,10,11") from None + indices = [index + count if index < 0 else index for index in requested] + if any(index < 0 or index >= count for index in indices): + raise ValueError(f"Layer index out of range: model has {count} layers/stages") + if len(set(indices)) != len(indices): + raise ValueError("Layer indices must not repeat") + return indices + + +def read_model_config(model_dir): + path = Path(model_dir) / "config.json" + if not path.exists(): + return {} + config = json.loads(path.read_text(encoding="utf-8-sig")) + if not isinstance(config, dict): + raise ValueError(f"{path} must contain a JSON object") + return config + + +def find_checkpoint(model_dir): + directory = Path(model_dir) + config = read_model_config(directory) + if config.get("checkpoint"): + path = directory / config["checkpoint"] + if not path.is_file(): + raise FileNotFoundError(path) + return path + for name in ("best.pt", "best_checkpoint.pth", "model.safetensors", "pytorch_model.bin"): + if (directory / name).is_file(): + return directory / name + paths = sorted(p for p in directory.iterdir() if p.suffix.lower() in CHECKPOINT_EXTENSIONS) + if len(paths) != 1: + raise ValueError(f"Expected one checkpoint in {directory}; set 'checkpoint' in config.json to select one") + return paths[0] + + +def model_folders(models_dir): + """Discover model subfolders under models/kaloscope.""" + root = Path(models_dir) / 'kaloscope' + return {path.name: path for path in sorted(root.iterdir()) if path.is_dir()} if root.is_dir() else {} + + +def load_checkpoint_payload(path): + if Path(path).suffix.lower() == ".safetensors": + from safetensors.torch import load_file + return load_file(str(path), device="cpu") + # Local training checkpoints also contain optimizer/RNG state (including numpy). + return torch.load(path, map_location="cpu", weights_only=False) + + +def normalize_state_dict_keys(state): + result = {} + for key, value in state.items(): + while key.startswith(("module.", "_orig_mod.")): + key = key.split(".", 1)[1] + result[key] = value + return result + + +def checkpoint_state(payload): + if not isinstance(payload, dict): + raise ValueError("Checkpoint must contain a state dictionary") + for key in ("model", "state_dict", "model_ema", "teacher", "student"): + if isinstance(payload.get(key), dict): + return checkpoint_state(payload[key]) + state = {k: v for k, v in payload.items() if isinstance(v, torch.Tensor)} + if not state: + raise ValueError("Checkpoint contains no model tensors") + return normalize_state_dict_keys(state) + + +def load_checkpoint_state(path): + return checkpoint_state(load_checkpoint_payload(path)) + + +def load_class_mapping(path): + if not path: + return None + with Path(path).open(encoding="utf-8-sig", newline="") as stream: + reader = csv.DictReader(stream) + if not reader.fieldnames or not {"class_id", "class_name"}.issubset(reader.fieldnames): + raise ValueError("CSV must contain class_id and class_name columns") + mapping = {int(row["class_id"]): row["class_name"] for row in reader} + if not mapping: + raise ValueError("Class mapping is empty") + return mapping + + +class DinoInferenceModel(nn.Module): + """Expose the same return_features interface as LSNet without random heads.""" + def __init__(self, backbone, pooling, head=None, projector=None, feature_source="backbone"): + super().__init__() + self.backbone = backbone + self.pooling = pooling + self.head = head + self.projector = projector + self.feature_source = feature_source + self.has_classifier = head is not None + self.pooled_dim = backbone.embed_dim * (2 if pooling == "cls_mean" else 1) + self.feature_dim = projector[-1].out_features if feature_source == "projector" else self.pooled_dim + + @staticmethod + def _pool(cls, patches, pooling): + if pooling == "cls": + return cls + if pooling == "mean": + return patches.mean(1) + return torch.cat((cls, patches.mean(1)), dim=1) + + def _intermediate(self, images, output_type, layers, norm): + indices = _selected_layers(layers, self.backbone.n_blocks) + kind = output_type.removeprefix("intermediate_") + is_vit = hasattr(self.backbone, "blocks") + if kind == "storage_tokens" and not self.backbone.n_storage_tokens: + raise ValueError("This architecture has no storage/register tokens") + kwargs = dict(n=sorted(indices), reshape=kind == "patch_map", return_class_token=True, + norm=False if kind == "prenorm" else norm) + if is_vit: + kwargs["return_extra_tokens"] = True + outputs = self.backbone.get_intermediate_layers(images, **kwargs) + selected = {} + for index, result in zip(sorted(indices), outputs): + patches, cls = result[:2] + storage = result[2] if is_vit else cls.new_empty(cls.shape[0], 0, cls.shape[-1]) + if kind in ("cls", "mean", "cls_mean"): + tensor = self._pool(cls, patches, kind) + elif kind in ("patch_tokens", "patch_map"): + tensor = patches + elif kind == "storage_tokens": + tensor = storage + elif kind == "all_tokens" or (kind == "prenorm" and is_vit): + tensor = torch.cat((cls.unsqueeze(1), storage, patches), dim=1) + elif kind == "prenorm": + tensor = patches + else: + raise ValueError(f"Unsupported intermediate output: {kind}") + selected[index] = tensor + if len({tuple(value.shape) for value in selected.values()}) != 1: + raise ValueError("Selected stages have different tensor shapes; select one ConvNeXt stage at a time") + return torch.stack([selected[index] for index in indices], dim=1) + + @torch.inference_mode() + def extract_tensor(self, images, output_type="default", layers="-1", intermediate_norm=True): + if output_type not in FEATURE_OUTPUTS: + raise ValueError(f"Unsupported feature output: {output_type}") + if output_type.startswith("intermediate_"): + return self._intermediate(images, output_type, layers, intermediate_norm) + if output_type == "patch_map": + return self._intermediate(images, "intermediate_patch_map", "-1", True)[:, 0] + tokens = self.backbone.forward_features(images) + cls, patches = tokens["x_norm_clstoken"], tokens["x_norm_patchtokens"] + if output_type in ("cls", "mean", "cls_mean"): + return self._pool(cls, patches, output_type) + if output_type == "patch_tokens": + return patches + if output_type == "storage_tokens": + if not self.backbone.n_storage_tokens: + raise ValueError("This architecture has no storage/register tokens") + return tokens["x_storage_tokens"] + if output_type == "all_tokens": + return torch.cat((cls.unsqueeze(1), tokens["x_storage_tokens"], patches), dim=1) + if output_type == "prenorm": + return tokens["x_prenorm"] + pooled = self._pool(cls, patches, self.pooling) + if output_type == "projector" or (output_type == "default" and self.feature_source == "projector"): + if self.projector is None: + raise ValueError("This checkpoint has no supported projector") + return self.projector(pooled) + return pooled + + def forward(self, images, return_features=False, return_both=False): + tokens = self.backbone.forward_features(images) + cls = tokens["x_norm_clstoken"] + patches = tokens["x_norm_patchtokens"] + features = self._pool(cls, patches, self.pooling) + output = self.projector(features) if self.feature_source == "projector" else features + if return_features: + return output + if self.head is None: + raise ValueError("This checkpoint has no classification head; use feature extraction or similarity") + logits = self.head(features) + return (output, logits) if return_both else logits + + +def _load_dino(name, state, model_config, feature_source=None): + from kaloscope_dinov3.architecture import build_backbone + # Never load the training machine's weights path or download pretrained weights. + backbone = build_backbone({**model_config, "name": name}) + prefixed = any(k.startswith("backbone.") for k in state) + backbone_state = {k.removeprefix("backbone."): v for k, v in state.items() if k.startswith("backbone.")} if prefixed else { + k: v for k, v in state.items() + if not k.startswith(("head.", "linear_head.", "projector.")) and k not in ("log_temperature", "bias") + } + backbone.load_state_dict(backbone_state, strict=True) + head_state = None + for prefix in ("head.", "linear_head."): + if prefix + "weight" in state: + head_state = {k.removeprefix(prefix): v for k, v in state.items() if k.startswith(prefix)} + break + pooling = model_config.get("pooling") + inferred_dim = head_state["weight"].shape[1] if head_state else ( + state["projector.0.weight"].shape[1] if "projector.0.weight" in state else backbone.embed_dim + ) + pooling = pooling or ("cls_mean" if inferred_dim == 2 * backbone.embed_dim else "cls") + if pooling not in ("cls", "mean", "cls_mean"): + raise ValueError("pooling must be cls, mean or cls_mean") + dimension = backbone.embed_dim * (2 if pooling == "cls_mean" else 1) + head = None + if head_state is not None: + count, head_dim = head_state["weight"].shape + if head_dim != dimension: + raise ValueError("Classifier input dimension differs from pooling output") + head = nn.Linear(dimension, count, bias="bias" in head_state) + head.load_state_dict(head_state, strict=True) + temporal = "log_temperature" in state and "bias" in state + source = feature_source or model_config.get("feature_source") or ("projector" if temporal else "backbone") + if source not in ("backbone", "projector"): + raise ValueError("feature_source must be backbone or projector") + projector = None + if source == "projector" or "projector.0.weight" in state or "projector.2.weight" in state: + if "projector.2.weight" not in state or "projector.0.weight" not in state: + raise ValueError("Checkpoint has no supported projector") + hidden, input_dim = state["projector.0.weight"].shape + output_dim, output_hidden = state["projector.2.weight"].shape + if input_dim != dimension or output_hidden != hidden: + raise ValueError("Projector dimensions differ from pooling output") + projector = nn.Sequential(nn.Linear(dimension, hidden), nn.GELU(), nn.Linear(hidden, output_dim)) + projector.load_state_dict({k.removeprefix("projector."): v for k, v in state.items() if k.startswith("projector.")}, strict=True) + return DinoInferenceModel(backbone, pooling, head, projector, source), temporal + + +def _load_lsnet(name, state): + # DINOv3 does not import LSNet's Triton kernels or depend on timm registration. + from lsnet_model import lsnet_artist + from timm.models import create_model + weight = state.get("head.l.weight") + feature_dim = weight.shape[1] if weight is not None else None + if "projection.0.l.weight" in state: + feature_dim = state["projection.0.l.weight"].shape[0] + model = create_model(name, pretrained=False, num_classes=weight.shape[0] if weight is not None else 0, + feature_dim=feature_dim, distillation="head_dist.l.weight" in state) + model.load_state_dict(state, strict=True) + model.has_classifier = weight is not None + return model + + +def load_model_bundle(model_dir=None, device="cuda", checkpoint=None, model_name=None, + class_csv=None, input_size=None, feature_source=None): + path = Path(checkpoint) if checkpoint else find_checkpoint(model_dir) + config = read_model_config(model_dir or path.parent) + selection = config.get("model", model_name) + if selection is None: + raise ValueError("Set the model architecture in config.json ('model') or pass an explicit model_name") + payload = load_checkpoint_payload(path) + state = checkpoint_state(payload) + embedded = payload.get("model_config", {}) + embedded = embedded if isinstance(embedded, dict) else {} + if isinstance(selection, dict): + name = selection.get("name") + options = {**embedded, **selection} + else: + name = selection + options = {**embedded, **{k: config[k] for k in ("kwargs", "pooling", "feature_source") if k in config}} + if not isinstance(name, str): + raise ValueError("config.json 'model' must be an architecture name or object with 'name'") + if name.startswith("dinov3_") or name in ("custom_vit", "custom_convnext"): + if embedded.get("name") and embedded["name"] != name: + raise ValueError("config.json architecture differs from checkpoint model_config") + model, temporal = _load_dino(name, state, options, feature_source) + from kaloscope_dinov3.preprocessing import image_transform, prepare_rgb + from torchvision import transforms + training_config = payload.get("config", {}) + data = training_config.get("data", {}) if isinstance(training_config, dict) else {} + data = {**data, **config.get("data", {})} + size = config.get("input_size", data.get("image_size", input_size or (512 if temporal else 224))) + custom_transform = data.get("custom_transform") + if custom_transform and custom_transform.startswith("dinov3."): + custom_transform = custom_transform.replace("dinov3.", "kaloscope_dinov3.", 1) + if temporal and custom_transform is None: + transform = transforms.Compose([prepare_rgb, transforms.Resize(size), transforms.CenterCrop(size), + transforms.ToTensor(), transforms.Normalize(data.get("mean") or [0.485, 0.456, 0.406], + data.get("std") or [0.229, 0.224, 0.225])]) + else: + transform = image_transform(size, data.get("mean"), data.get("std"), custom_transform=custom_transform) + elif name in LSNET_MODELS: + model = _load_lsnet(name, state) + from kaloscope_dinov3.preprocessing import prepare_rgb + from torchvision import transforms + from lsnet_model.lsnet_artist import default_cfgs_artist + from timm.data import resolve_data_config, create_transform + size = config.get("input_size", input_size or default_cfgs_artist[name]["input_size"][1]) + transform = transforms.Compose([prepare_rgb, create_transform( + **resolve_data_config({"input_size": (3, size, size)}, model=model))]) + else: + raise ValueError(f"Unsupported model architecture in config.json: {name}") + csv_path = Path(class_csv) if class_csv else path.parent / "class_mapping.csv" + if class_csv and not csv_path.is_file(): + raise FileNotFoundError(csv_path) + mapping = load_class_mapping(csv_path) if csv_path.is_file() else None + classes = config.get("classes", payload.get("classes")) + if mapping is None and classes is not None: + mapping = {i: str(label) for i, label in enumerate(classes)} if isinstance(classes, list) else { + int(i): str(label) for i, label in classes.items() + } + count = model.head.out_features if isinstance(model.head, nn.Linear) else ( + model.head.l.out_features if model.has_classifier else 0) + if model.has_classifier and mapping is not None and set(mapping) != set(range(count)): + raise ValueError("Class mapping IDs must exactly match classifier outputs") + model.to(device).eval() + return {"model": model, "transform": transform, "class_mapping": mapping or {}, "device": device, + "model_type": name, "has_classifier": model.has_classifier, "feature_dim": model.feature_dim, + "feature_source": getattr(model, "feature_source", "backbone"), "input_size": size, + "checkpoint": str(path)} diff --git a/pyproject.toml b/pyproject.toml index b24b341..4c65f64 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] -name = "lsnet" -description = "" -version = "1.0.7" +name = "comfyui-kaloscope" +description = "LSNet and DINOv3 inference, features and analysis for ComfyUI, WebUI and standalone use" +version = "1.0.0" license = {file = "LICENSE"} # classifiers = [ # # For OS-independent nodes (works on all operating systems) @@ -20,7 +20,10 @@ license = {file = "LICENSE"} # "Environment :: GPU :: Apple Metal", # Apple Metal support # ] -dependencies = ["einops>=0.8.1", "fvcore", "easydict", "matplotlib", "yacs", "scikit-image", "wandb", "torch>=2.4.1", "torchvision>=0.11.0", "timm>=1.0.20", "numpy>=1.19.0", "Pillow>=8.0.0", "scikit-learn>=1.0.0", "matplotlib>=3.3.0", "tqdm>=4.60.0", "triton-windows"] +dependencies = ["einops>=0.8.1", "fvcore", "easydict", "matplotlib", "yacs", "scikit-image", "wandb", "torch>=2.4.1", "torchvision>=0.11.0", "timm>=1.0.20", "numpy>=1.19.0", "Pillow>=8.0.0", "scikit-learn>=1.2.0", "scipy>=1.8.0", "matplotlib>=3.3.0", "tqdm>=4.60.0", "safetensors>=0.4.0", "triton-windows; sys_platform == 'win32'"] + +[project.optional-dependencies] +webui = ["gradio>=3.41,<6", "fastapi", "uvicorn", "python-multipart"] [project.urls] Repository = "https://github.com/spawner1145/comfyui-lsnet" @@ -30,7 +33,7 @@ Documentation = "https://github.com/spawner1145/comfyui-lsnet/wiki" [tool.comfy] PublisherId = "spawner" -DisplayName = "comfyui-lsnet" +DisplayName = "comfyui-kaloscope" Icon = "" includes = [] # "requires-comfyui" = ">=1.0.0" # ComfyUI version compatibility diff --git a/requirements-webui.txt b/requirements-webui.txt new file mode 100644 index 0000000..a57b4fd --- /dev/null +++ b/requirements-webui.txt @@ -0,0 +1,5 @@ +-r requirements.txt +gradio>=3.41,<6 +fastapi +uvicorn +python-multipart diff --git a/requirements.txt b/requirements.txt index 18f2458..40b03fd 100644 --- a/requirements.txt +++ b/requirements.txt @@ -13,7 +13,8 @@ timm>=1.0.20 numpy>=1.19.0 Pillow>=8.0.0 -scikit-learn>=1.0.0 +scikit-learn>=1.2.0 +scipy>=1.8.0 matplotlib>=3.3.0 @@ -24,3 +25,6 @@ tqdm>=4.60.0 # Windows triton-windows; sys_platform == 'win32' + +# DINOv3 safetensors checkpoints +safetensors>=0.4.0 diff --git a/scripts/app.py b/scripts/app.py index d018e92..dc5a4d1 100644 --- a/scripts/app.py +++ b/scripts/app.py @@ -1,3 +1,11 @@ +import sys +from pathlib import Path + +# WebUI extensions and direct script execution do not put this plugin root on sys.path. +PLUGIN_ROOT = Path(__file__).resolve().parents[1] +if str(PLUGIN_ROOT) not in sys.path: + sys.path.insert(0, str(PLUGIN_ROOT)) + import logging import threading from threading import Lock @@ -13,9 +21,10 @@ import argparse logging.basicConfig(level=logging.INFO) def parse_args(): - parser = argparse.ArgumentParser(description="LSNet Artist Inference WebUI") + parser = argparse.ArgumentParser(description="Kaloscope Artist Inference WebUI") parser.add_argument("--host", type=str, default="127.0.0.1", help="Server host") parser.add_argument("--port", type=int, default=7860, help="Server port") + parser.add_argument("--models-dir", type=str, default=None, help="Root containing kaloscope/ model folders") return parser.parse_args() try: @@ -33,18 +42,21 @@ if IN_WEBUI: from backend_lsnet.api import on_app_started def on_ui_tabs(): block = create_ui() - return [(block, "LSNet Artist", "lsnet_tab")] + return [(block, "Kaloscope", "kaloscope_tab")] script_callbacks.on_ui_tabs(on_ui_tabs) script_callbacks.on_app_started(on_app_started) else: if __name__ == "__main__": args = parse_args() + if args.models_dir: + os.environ['KALOSCOPE_MODELS_DIR'] = str(Path(args.models_dir).resolve()) # Create models directory - os.makedirs("models/lsnet", exist_ok=True) + from backend_lsnet.model_paths import models_root + (models_root() / 'kaloscope').mkdir(parents=True, exist_ok=True) app = FastAPI(docs_url="/docs", openapi_url="/openapi.json") queue_lock = Lock() - api = Api(app, queue_lock, prefix="/lsnet/v1") + api = Api(app, queue_lock, prefix="/kaloscope/v1") logging.info("API 路由已挂载到 FastAPI 实例") block = create_ui() @@ -64,4 +76,4 @@ else: log_level="info" ) except Exception as e: - logging.error(f"启动失败: {str(e)}") \ No newline at end of file + logging.error(f"启动失败: {str(e)}") diff --git a/tests/render_analysis_examples.py b/tests/render_analysis_examples.py new file mode 100644 index 0000000..5790082 --- /dev/null +++ b/tests/render_analysis_examples.py @@ -0,0 +1,145 @@ +"""Infer each example once, cache features, then exercise every plotting node mode.""" +import argparse +import csv +import hashlib +import html +import json +from pathlib import Path +import re +import sys + +import numpy as np +from PIL import Image, ImageDraw, ImageFont, ImageOps +import torch + +ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(ROOT)) +from feature_analysis import CHART_TYPES +from kaloscope_dinov3.preprocessing import prepare_rgb +from model_loading import load_model_bundle +from test_model_loading import load_nodes + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument('--input', default='../test') + parser.add_argument('--model-dir', default='../model') + parser.add_argument('--output', default='../outputs') + parser.add_argument('--device', default='cuda') + parser.add_argument('--batch-size', default=2, type=int) + parser.add_argument('--reuse-cache', action='store_true', help='Redraw charts from features.pt without loading the model') + args = parser.parse_args() + torch.set_num_threads(4) + source = Path(args.input).resolve() + output = Path(args.output).resolve() + output.mkdir(parents=True, exist_ok=True) + paths = sorted(path for path in source.iterdir() if path.suffix.lower() in {'.png','.jpg','.jpeg','.webp','.bmp'}) + if not paths: + raise ValueError(f'No images in {source}') + if args.batch_size < 1: + raise ValueError('batch-size must be positive') + cache_path=output/'features.pt' + if args.reuse_cache: + cached=torch.load(cache_path,map_location='cpu',weights_only=True) + features,patch_features=cached['features'],cached['patch_tokens'] + labels,records=cached['labels'],cached['source_images'] + model_info=cached.get('model_info',{'model_type':'dinov3_vitb16','feature_source':'projector', + 'input_size':512,'checkpoint':str(Path(args.model_dir).resolve()/'best.pt')}) + inference_batches=cached.get('inference_batches',4) + previews=[] + for record in records: + with Image.open(record['path']) as original: + previews.append(ImageOps.pad(prepare_rgb(original),(240,200),color='white')) + thumbnails=torch.stack([torch.from_numpy(np.array(image)).float()/255 for image in previews]) + print('Reusing cached features: zero new model forwards',flush=True) + else: + bundle = load_model_bundle(args.model_dir, device=args.device) + model = bundle['model'] + model_info = {key: bundle[key] for key in ('model_type','feature_source','input_size','checkpoint')} + if not hasattr(model, 'extract_tensor'): + raise ValueError('This example suite requires a DINOv3 checkpoint for patch analysis') + vectors, patches, previews = [], [], [] + labels, records = [], [] + inference_batches = 0 + with torch.inference_mode(): + for start in range(0,len(paths),args.batch_size): + batch = [] + for index, path in enumerate(paths[start:start+args.batch_size], start): + with Image.open(path) as original: + image = prepare_rgb(original) + batch.append(bundle['transform'](image)) + previews.append(ImageOps.pad(image, (240,200), color='white')) + artist = re.search(r'_drawn_by_(.+?)__[0-9a-f]+$', path.stem) + label = f'{index+1:02d} ' + (artist.group(1) if artist else path.stem[:22]) + labels.append(label) + with path.open('rb') as stream: + digest = hashlib.file_digest(stream, 'sha256').hexdigest() + records.append({'index': index, 'label': label, 'path': str(path), 'sha256': digest}) + tokens = model.backbone.forward_features(torch.stack(batch).to(args.device)) + pooled = model._pool(tokens['x_norm_clstoken'],tokens['x_norm_patchtokens'],model.pooling) + feature = model.projector(pooled) if model.feature_source=='projector' else pooled + vectors.append(feature.float().cpu()) + patches.append(tokens['x_norm_patchtokens'].float().cpu()) + inference_batches += 1 + print(f'Inferred images {start+1}-{min(start+args.batch_size,len(paths))} once',flush=True) + features = torch.cat(vectors) + patch_features = torch.cat(patches) + thumbnails = torch.stack([torch.from_numpy(np.array(image)).float()/255 for image in previews]) + torch.save({'features':features, 'patch_tokens':patch_features, 'labels':labels, 'source_images':records,'model_info':model_info,'inference_batches':inference_batches},output/'features.pt') + # Release GPU model; all charts below consume cached CPU tensors only. + del model, bundle + if torch.cuda.is_available(): + torch.cuda.empty_cache() + nodes = load_nodes() + node = nodes.KaloscopeFeatureAnalysisNode() + generated = [] + for chart in CHART_TYPES: + tensor = patch_features if chart=='patch_energy' else features + result, text, distances = node.analyze(tensor,chart_type=chart,labels=json.dumps(labels),images=thumbnails, + tensor_layout='tokens' if chart=='patch_energy' else 'vectors',n_clusters=3,top_k=3, + metric='cosine',normalize=True,width=1600,height=1100,perplexity=3.0) + report = json.loads(text) + report['source_images']=records + report['feature_source']='patch_tokens' if chart=='patch_energy' else model_info['feature_source'] + array = (result[0].clamp(0,1).numpy()*255).round().astype(np.uint8) + Image.fromarray(array).save(output/f'{chart}.png',dpi=(120,120)) + (output/f'{chart}.json').write_text(json.dumps(report,ensure_ascii=False,indent=2,allow_nan=False),encoding='utf-8') + if chart=='relationship_graph': + with (output/'distance_matrix.csv').open('w',encoding='utf-8-sig',newline='') as stream: + writer=csv.writer(stream) + writer.writerow(['image',*labels]) + writer.writerows([label,*row.tolist()] for label,row in zip(labels,distances)) + generated.append({'type':chart,'image':f'{chart}.png','report':f'{chart}.json'}) + print(f'Saved {chart}.png',flush=True) + manifest={'input_directory':str(source),'checkpoint':model_info['checkpoint'], + 'architecture':model_info['model_type'],'model_info':model_info,'feature_shape':list(features.shape),'patch_shape':list(patch_features.shape), + 'source_images':records,'inference_batches':inference_batches,'image_forward_count':len(paths), + 'analysis_model_forward_count':0,'this_run_model_forward_count':0 if args.reuse_cache else len(paths),'seed':42,'charts':generated} + (output/'manifest.json').write_text(json.dumps(manifest,ensure_ascii=False,indent=2),encoding='utf-8') + body=''.join(f'

{html.escape(item["type"])}

{html.escape(item[

分析数据 JSON

' for item in generated) + (output/'index.html').write_text('''Kaloscope 分析图示例 + +

Kaloscope 特征分析示例

每张输入图片仅推理一次。18 类图表复用缓存特征;二维距离布局是近似结果,具体距离以矩阵和 JSON 为准。

+

来源与运行记录 · 距离矩阵 CSV

'''+body+'',encoding='utf-8') + contact = Image.new('RGB',(1600,math_rows(len(generated))*390),'#eef2f5') + drawer=ImageDraw.Draw(contact) + try: + font=ImageFont.truetype('C:/Windows/Fonts/arial.ttf',22) + except OSError: + font=ImageFont.load_default() + for index,item in enumerate(generated): + x=(index%3)*530+10; y=(index//3)*390 + with Image.open(output/item['image']) as image: + preview=ImageOps.contain(image,(510,350)) + contact.paste(preview,(x,y+32)) + drawer.text((x,y+5),item['type'],fill='#243746',font=font) + contact.save(output/'overview.png') + print(f'{len(generated)} chart types written to {output}; analysis performed zero model forwards',flush=True) + + +def math_rows(count): + return (count+2)//3 + + +if __name__=='__main__': + main() diff --git a/tests/smoke_checkpoint.py b/tests/smoke_checkpoint.py new file mode 100644 index 0000000..5478817 --- /dev/null +++ b/tests/smoke_checkpoint.py @@ -0,0 +1,124 @@ +"""Compare a supplied temporal checkpoint against its original training encoder. + +Run from the plugin directory: +python tests/smoke_checkpoint.py --model-dir ../model --source ../kaloscope-dinov3 +""" +import argparse +import json +from pathlib import Path +import sys +import tempfile +import types + +import numpy as np +from PIL import Image +import torch +from torchvision import transforms + +ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(ROOT)) +from model_loading import load_model_bundle, load_checkpoint_payload +from test_model_loading import load_nodes + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--model-dir", required=True) + parser.add_argument("--source", required=True) + parser.add_argument("--device", default="cuda") + args = parser.parse_args() + sys.path.insert(0, str(Path(args.source).resolve())) + from dinov3.finetune.temporal.model import Encoder + bundle = load_model_bundle(args.model_dir, device=args.device) + payload = load_checkpoint_payload(bundle["checkpoint"]) + original = Encoder(payload["model_config"], initialize=False) + original.load_state_dict(payload["model"], strict=True) + original.to(args.device).eval() + image = Image.fromarray(np.random.default_rng(42).integers(0, 256, (480, 640, 3), dtype=np.uint8)) + spatial = transforms.Compose([transforms.Resize(512), transforms.CenterCrop(512), transforms.PILToTensor()]) + raw = spatial(image).unsqueeze(0).to(args.device) + batch = bundle["transform"](image).unsqueeze(0).to(args.device) + with torch.inference_mode(): + expected = original(raw)[:, original.feature_dim:] + actual = bundle["model"](batch, return_features=True) + torch.testing.assert_close(actual, expected, rtol=1e-5, atol=1e-5) + assert torch.isfinite(actual).all() + print(json.dumps({"architecture": bundle["model_type"], "has_classifier": bundle["has_classifier"], + "input_size": bundle["input_size"], "feature_source": bundle["feature_source"], + "features": list(actual.shape), "max_error_vs_training_encoder": (actual - expected).abs().max().item()}, indent=2), flush=True) + with torch.inference_mode(): + tokens = original.backbone.forward_features(batch) + cls, patches, storage = tokens['x_norm_clstoken'], tokens['x_norm_patchtokens'], tokens['x_storage_tokens'] + pooled = torch.cat((cls, patches.mean(1)), dim=1) + references = { + 'default': actual, 'backbone': pooled, 'cls': cls, 'mean': patches.mean(1), + 'cls_mean': pooled, 'projector': actual, 'patch_tokens': patches, + 'patch_map': patches.transpose(1, 2).reshape(1, 768, 32, 32), + 'storage_tokens': storage, 'all_tokens': torch.cat((cls.unsqueeze(1), storage, patches), dim=1), + 'prenorm': tokens['x_prenorm'], + } + shapes = {} + for kind, reference in references.items(): + result = bundle['model'].extract_tensor(batch, kind) + torch.testing.assert_close(result, reference, rtol=1e-5, atol=1e-5) + shapes[kind] = list(result.shape) + native = original.backbone.get_intermediate_layers(batch, n=[8, 9, 10, 11], + return_class_token=True, return_extra_tokens=True) + for kind in ('cls', 'mean', 'cls_mean', 'patch_tokens', 'patch_map', 'storage_tokens', 'all_tokens'): + expected_layers = [] + for patch, cls_layer, registers in native: + if kind == 'cls': + value = cls_layer + elif kind == 'mean': + value = patch.mean(1) + elif kind == 'cls_mean': + value = torch.cat((cls_layer, patch.mean(1)), dim=1) + elif kind == 'patch_tokens': + value = patch + elif kind == 'patch_map': + value = patch.transpose(1, 2).reshape(1, 768, 32, 32) + elif kind == 'storage_tokens': + value = registers + else: + value = torch.cat((cls_layer.unsqueeze(1), registers, patch), dim=1) + expected_layers.append(value) + output_type = 'intermediate_' + kind + result = bundle['model'].extract_tensor(batch, output_type, '-4,-3,-2,-1') + torch.testing.assert_close(result, torch.stack(expected_layers, dim=1), rtol=1e-5, atol=1e-5) + shapes[output_type] = list(result.shape) + native_raw = original.backbone.get_intermediate_layers(batch, n=[8, 9, 10, 11], norm=False, + return_class_token=True, return_extra_tokens=True) + expected_raw = torch.stack([torch.cat((cls_layer.unsqueeze(1), registers, patch), dim=1) + for patch, cls_layer, registers in native_raw], dim=1) + result = bundle['model'].extract_tensor(batch, 'intermediate_prenorm', '8,9,10,11') + torch.testing.assert_close(result, expected_raw, rtol=1e-5, atol=1e-5) + shapes['intermediate_prenorm'] = list(result.shape) + print('All feature outputs match original backbone:', json.dumps(shapes), flush=True) + del original, payload + nodes = load_nodes() + image_tensor = torch.from_numpy(np.array(image)).float().unsqueeze(0) / 255 + tags, features_json = nodes.KaloscopeArtistInferenceNode().process(image_tensor, bundle, 5, 0.0) + assert tags == "" and len(json.loads(features_json)["features"]) == 256 + features = nodes.KaloscopeExtractFeaturesNode().extract(image_tensor, bundle)[0] + torch.testing.assert_close(features, actual.cpu(), rtol=1e-4, atol=1e-4) + similarity = json.loads(nodes.KaloscopeArtistSimilarityNode().process(image_tensor, image_tensor, bundle)[0]) + assert abs(similarity["similarities"][0] - 1.0) < 1e-5 + print("ComfyUI inference, extraction and self-similarity passed", flush=True) + from backend_lsnet.inference import process_image_from_pil + output = process_image_from_pil(image, checkpoint=bundle["checkpoint"], device=args.device) + np.testing.assert_allclose(output["features"], actual[0].cpu().numpy(), rtol=1e-4, atol=1e-4) + print("Standalone/API backend output matches", flush=True) + from inference_artist import get_args_parser, main as cli_main + with tempfile.TemporaryDirectory() as temp: + path = Path(temp) + image.save(path / "input.png") + cli_args = get_args_parser().parse_args(["--checkpoint", bundle["checkpoint"], "--input", str(path / "input.png"), + "--output", str(path / "output"), "--device", args.device]) + cli_main(cli_args) + result = json.loads((path / "output/input_result.json").read_text(encoding="utf-8")) + np.testing.assert_allclose(result["features"], actual[0].cpu().numpy(), rtol=1e-4, atol=1e-4) + print("CLI auto inference matches", flush=True) + + +if __name__ == "__main__": + main() diff --git a/tests/smoke_entry_points.py b/tests/smoke_entry_points.py new file mode 100644 index 0000000..7f2ef4b --- /dev/null +++ b/tests/smoke_entry_points.py @@ -0,0 +1,99 @@ +"""Check supplied images/checkpoint and a live standalone Gradio/API server.""" +import argparse +import base64 +import json +from pathlib import Path +import shutil +import sys +import urllib.request +from unittest.mock import patch + +ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(ROOT)) +import numpy as np +import torch +from PIL import Image +from gradio_client import Client, handle_file + +from backend_lsnet.analysis import cache_bytes, read_cache +from backend_lsnet.analysis_api import AnalysisOptions, values +from backend_lsnet.analysis_ui import extract_uploaded, ANALYSIS_OPTION_NAMES +from feature_analysis import CHART_TYPES +from model_loading import find_checkpoint +from analysis_cli import parser as cli_parser, main as cli_main + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--input', type=Path, default=Path('../test')) + parser.add_argument('--model-dir', type=Path, default=Path('../model')) + parser.add_argument('--output', type=Path, default=Path('../outputs/entry_points')) + parser.add_argument('--patch-cache', type=Path, default=Path('../outputs/features.pt')) + parser.add_argument('--server-url', default='http://127.0.0.1:7865') + parser.add_argument('--device', default='cuda') + args = parser.parse_args() + torch.set_num_threads(4) + args.output.mkdir(parents=True, exist_ok=True) + paths = sorted(path for path in args.input.iterdir() if path.suffix.lower() in {'.png', '.jpg', '.jpeg', '.webp', '.bmp'}) + # Test actual extraction with the provided checkpoint; replace only folder discovery. + with patch('backend_lsnet.analysis_ui.get_checkpoint_path', return_value=str(find_checkpoint(args.model_dir))): + cached, status, downloaded = extract_uploaded([str(path.resolve()) for path in paths], 'provided', args.device, 'default', '-1', True, 2) + (args.output / 'features.npz').write_bytes(Path(downloaded).read_bytes()) + Path(downloaded).unlink() + Path(downloaded).parent.rmdir() + print(status, flush=True) + original = torch.load(args.patch_cache, map_location='cpu', weights_only=True) + torch.testing.assert_close(cached['features'], original['features'], rtol=1e-4, atol=1e-4) + # Reuse previously inferred patches to test every chart without extra inference. + patch_npz = args.output / 'patch_features.npz' + patch_npz.write_bytes(cache_bytes(original['patch_tokens'], cached['labels'], 'patch_tokens')) + payload = base64.b64encode(patch_npz.read_bytes()).decode() + client = Client(args.server_url, verbose=False) + imported = client.predict(handle_file(str(patch_npz.resolve())), api_name='/import_uploaded_cache') + print(imported, flush=True) + options = values(AnalysisOptions(width=1200, height=900, perplexity=3, top_k=3)) + options['labels'] = '\n'.join(original['labels']) + api_options = {**options, 'labels': original['labels']} + recorded = [] + for chart in CHART_TYPES: + ui_dir = args.output / 'ui' + ui_dir.mkdir(exist_ok=True) + result = client.predict(chart, *[options[name] for name in ANALYSIS_OPTION_NAMES], api_name='/plot_uploaded_cache') + shutil.copyfile(result[0], ui_dir / f'{chart}.png') + (ui_dir / f'{chart}.json').write_text(result[1], encoding='utf-8') + for source in result[2]: + if str(source).endswith('.csv'): + shutil.copyfile(source, ui_dir / f'{chart}.csv') + request = urllib.request.Request(args.server_url + '/kaloscope/v1/analyze', + data=json.dumps({'cache_base64': payload, 'chart_type': chart, 'options': api_options}).encode(), + headers={'Content-Type': 'application/json'}) + with urllib.request.urlopen(request, timeout=60) as response: + api = json.load(response) + api_dir = args.output / 'api' + api_dir.mkdir(exist_ok=True) + (api_dir / f'{chart}.png').write_bytes(base64.b64decode(api['image_base64'])) + (api_dir / f'{chart}.json').write_text(json.dumps(api['analysis'], indent=2, ensure_ascii=False), encoding='utf-8') + np.testing.assert_allclose(api['distance_matrix'], json.loads(result[1])['distances'], atol=1e-7) + for directory in (ui_dir, api_dir): + with Image.open(directory / f'{chart}.png') as image: + assert image.size == (1200, 900) + image.verify() + recorded.append(chart) + print(f'Live UI and API cache rendering passed: {chart}', flush=True) + (args.output / 'options.json').write_text(json.dumps(api_options), encoding='utf-8') + cli_main(cli_parser().parse_args(['--features', str(patch_npz), '--all-charts', '--output', str(args.output / 'cli'), + '--options-json', str(args.output / 'options.json')])) + for chart in CHART_TYPES: + cli = json.loads((args.output / 'cli' / f'{chart}.json').read_text(encoding='utf-8')) + api = json.loads((args.output / 'api' / f'{chart}.json').read_text(encoding='utf-8')) + np.testing.assert_allclose(cli['distances'], api['distances'], atol=1e-7) + manifest = {'sample_count': len(paths), 'feature_shape': list(cached['features'].shape), + 'patch_shape': list(original['patch_tokens'].shape), 'charts_per_entry': recorded, + 'live_gradio_cache_session': True, 'analysis_model_forward_count': 0, + 'numeric_parity': 'UI / HTTP API / standalone CLI distances equal within 1e-7'} + (args.output / 'manifest.json').write_text(json.dumps(manifest, indent=2), encoding='utf-8') + print('All live entry-point checks passed', flush=True) + + +if __name__ == '__main__': + main() diff --git a/tests/test_entry_points.py b/tests/test_entry_points.py new file mode 100644 index 0000000..4add13d --- /dev/null +++ b/tests/test_entry_points.py @@ -0,0 +1,193 @@ +"""Check real UI/API/CLI adapters, numerical parity and cache-only rendering.""" +import base64 +import importlib.util +import io +import json +import os +from pathlib import Path +import tempfile +import types +import unittest +from unittest.mock import patch + +import numpy as np +from PIL import Image +import torch +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from test_model_loading import ReferenceInferenceModel, load_nodes, ROOT +from backend_lsnet.analysis import read_cache, cache_bytes, feature_tools +from backend_lsnet.analysis_api import AnalysisOptions, values +from backend_lsnet.analysis_ui import (extract_uploaded, import_uploaded_cache, plot_uploaded_cache, + ANALYSIS_OPTION_NAMES) +from backend_lsnet.api import Api +from backend_lsnet.ui import create_ui +from feature_analysis import CHART_TYPES +from model_loading import FEATURE_OUTPUTS, load_model_bundle +from analysis_cli import parser, main + + +class EntryPointTests(unittest.TestCase): + @classmethod + def setUpClass(cls): + torch.set_num_threads(2) + cls.temporary = tempfile.TemporaryDirectory() + cls.root = Path(cls.temporary.name) + cls.model_dir = cls.root / 'models' / 'kaloscope' / 'tiny' + cls.model_dir.mkdir(parents=True) + options = {'name': 'custom_vit', 'pooling': 'cls_mean', 'projection_dim': 8, + 'kwargs': {'embed_dim': 24, 'depth': 2, 'num_heads': 3, 'patch_size': 16, + 'n_storage_tokens': 2, 'pos_embed_rope_dtype': 'fp32'}} + torch.manual_seed(7) + model = ReferenceInferenceModel(options, 3).eval() + torch.save({'model': model.state_dict(), 'model_config': options, 'classes': ['a', 'b', 'c']}, cls.model_dir / 'best.pt') + (cls.model_dir / 'config.json').write_text(json.dumps({'model': 'custom_vit', 'input_size': 32}), encoding='utf-8') + cls.input_dir = cls.root / 'images' + cls.input_dir.mkdir() + cls.images = [Image.new('RGB', (48, 40), color) for color in ((220,20,60),(160,10,20),(20,180,80),(10,110,50),(50,40,220),(30,20,150))] + cls.files, cls.encoded = [], [] + for index, image in enumerate(cls.images): + path = cls.input_dir / f'{index}.png' + image.save(path) + cls.files.append(str(path)) + cls.encoded.append(base64.b64encode(path.read_bytes()).decode()) + cls.bundle = load_model_bundle(cls.model_dir, device='cpu') + + @classmethod + def tearDownClass(cls): + cls.temporary.cleanup() + + def setUp(self): + self.environment = patch.dict(os.environ, {'KALOSCOPE_MODELS_DIR': str(self.root / 'models')}) + self.environment.start() + + def tearDown(self): + self.environment.stop() + + def test_all_feature_outputs_are_identical_in_comfy_ui_and_api(self): + app = FastAPI() + api = Api(app) + nodes = load_nodes() + images = torch.stack([torch.from_numpy(np.array(image.resize((48, 40)))).float() / 255 for image in self.images]) + try: + with patch('backend_lsnet.analysis_ui.load_model_bundle', return_value=self.bundle), \ + patch('backend_lsnet.analysis_api.load_model_bundle', return_value=self.bundle), TestClient(app) as client: + for output_type in FEATURE_OUTPUTS: + with self.subTest(output_type=output_type): + layers = '-1,0' if output_type.startswith('intermediate_') else '-1' + expected = nodes.KaloscopeExtractFeaturesNode().extract(images, self.bundle, output_type, layers)[0] + cached, _, download = extract_uploaded(self.files, 'tiny', 'cpu', output_type, layers, True, 2) + torch.testing.assert_close(cached['features'], expected) + imported, _ = import_uploaded_cache(download) + torch.testing.assert_close(imported['features'], expected) + response = client.post('/kaloscope/v1/features', json={'input_images': self.encoded, + 'model_name': 'tiny', 'device': 'cpu', 'output_type': output_type, 'layers': layers}) + self.assertEqual(response.status_code, 200, response.text) + cache = read_cache(io.BytesIO(base64.b64decode(response.json()['cache_base64']))) + torch.testing.assert_close(cache['features'], expected) + Path(download).unlink() + Path(download).parent.rmdir() + finally: + api.executor.shutdown() + + def test_all_charts_work_through_ui_and_api_without_loading_a_model(self): + app = FastAPI() + api = Api(app) + from backend_lsnet.analysis import extract_batch + features = extract_batch(self.images, self.bundle, 'patch_tokens') + cached = {'features': features, 'labels': [f'Image {i}' for i in range(6)], 'output_type': 'patch_tokens', 'layers': '-1'} + encoded = base64.b64encode(cache_bytes(features, cached['labels'], 'patch_tokens')).decode() + options = values(AnalysisOptions(width=800, height=600, perplexity=3)) + try: + with patch('backend_lsnet.analysis_api.load_model_bundle', side_effect=AssertionError('must not infer')), \ + patch('backend_lsnet.analysis_ui.load_model_bundle', side_effect=AssertionError('must not infer')), TestClient(app) as client: + for chart in CHART_TYPES: + with self.subTest(chart=chart): + image, report, files = plot_uploaded_cache(cached, chart, *[options[name] for name in ANALYSIS_OPTION_NAMES]) + self.assertEqual(image.shape, (600, 800, 3)) + response = client.post('/kaloscope/v1/analyze', json={'cache_base64': encoded, 'chart_type': chart, 'options': options}) + self.assertEqual(response.status_code, 200, response.text) + result = response.json() + self.assertEqual(result['analysis']['chart_type'], chart) + np.testing.assert_allclose(result['distance_matrix'], json.loads(report)['distances'], atol=1e-7) + rendered = Image.open(io.BytesIO(base64.b64decode(result['image_base64']))) + self.assertEqual(rendered.size, (800, 600)) + for file in files: + Path(file).unlink() + Path(files[0]).parent.rmdir() + finally: + api.executor.shutdown() + + def test_image_api_extracts_once_and_reuses_returned_cache(self): + app = FastAPI() + api = Api(app) + try: + with patch('backend_lsnet.analysis_api.load_model_bundle', return_value=self.bundle) as loading, \ + patch.object(self.bundle['model'].backbone, 'forward_features', wraps=self.bundle['model'].backbone.forward_features) as forward, TestClient(app) as client: + response = client.post('/kaloscope/v1/analyze', json={'image_batch': {'input_images': self.encoded, + 'model_name': 'tiny', 'device': 'cpu', 'batch_size': 2}, 'options': {'width': 800, 'height': 600}}) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(loading.call_count, 1) + self.assertEqual(forward.call_count, 3) + result = client.post('/kaloscope/v1/analyze', json={'cache_base64': response.json()['cache_base64'], + 'chart_type': 'distance_heatmap', 'options': {'width': 800, 'height': 600}}) + self.assertEqual(result.status_code, 200, result.text) + self.assertEqual(forward.call_count, 3) + finally: + api.executor.shutdown() + + def test_feature_tools_known_values_and_http_input_errors(self): + features = torch.tensor([[1.,0.],[1.,1.],[0.,1.]]) + common = feature_tools(features) + np.testing.assert_allclose(common['common_features'], [2/3, 2/3]) + self.assertEqual(feature_tools(features, 'compare_groups', 0, ['a','a','b'])['best_group'], 'a') + app = FastAPI() + api = Api(app) + try: + with TestClient(app) as client: + result = client.post('/kaloscope/v1/feature-tools', json={'features': features.tolist(), 'operation': 'similarity'}) + self.assertEqual(result.status_code, 200) + np.testing.assert_allclose(result.json()['similarities'], [1., 1/np.sqrt(2), 0.]) + result = client.post('/kaloscope/v1/analyze', json={'features': features.tolist(), 'cache_base64': 'invalid'}) + self.assertEqual(result.status_code, 400) + result = client.post('/kaloscope/v1/analyze', json={'features': features.tolist(), 'chart_type': 'bad'}) + self.assertEqual(result.status_code, 400) + self.assertEqual(len(client.get('/kaloscope/v1/models').json()['chart_types']), 18) + finally: + api.executor.shutdown() + + def test_cli_cache_and_real_image_path_match(self): + output = self.root / 'cli' + main(parser().parse_args(['--input', str(self.input_dir), '--model-dir', str(self.model_dir), '--device', 'cpu', + '--output-type', 'intermediate_patch_map', '--layers=-1,0', '--output', str(output), + '--chart-type', 'patch_energy', '--width', '800', '--height', '600'])) + cached = read_cache(output / 'features.npz') + self.assertEqual(tuple(cached['features'].shape), (6, 2, 24, 2, 2)) + with patch('analysis_cli.load_model_bundle', side_effect=AssertionError('must not infer')): + main(parser().parse_args(['--features', str(output / 'features.npz'), '--all-charts', '--output', str(output / 'redraw'), + '--width', '800', '--height', '600', '--perplexity', '3'])) + self.assertEqual(len(list((output / 'redraw').glob('*.png'))), 18) + + def test_same_gradio_tabs_are_registered_by_webui_extension(self): + ui = create_ui() + callback_names = {callback.fn.__name__ for callback in ui.fns.values()} + self.assertTrue({'infer', 'extract_uploaded', 'import_uploaded_cache', 'plot_uploaded_cache', 'tools_uploaded_cache'} <= callback_names) + registrations = {} + callbacks = types.SimpleNamespace(on_ui_tabs=lambda callback: registrations.update(tabs=callback), + on_app_started=lambda callback: registrations.update(api=callback)) + modules = types.ModuleType('modules') + modules.script_callbacks = callbacks + modules.shared = types.SimpleNamespace() + spec = importlib.util.spec_from_file_location('test_webui_extension', ROOT / 'scripts' / 'app.py') + module = importlib.util.module_from_spec(spec) + with patch.dict('sys.modules', {'modules': modules}): + spec.loader.exec_module(module) + tabs = registrations['tabs']() + self.assertTrue(module.IN_WEBUI) + self.assertEqual(tabs[0][1:], ('Kaloscope', 'kaloscope_tab')) + self.assertIsNotNone(registrations['api']) + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/test_feature_analysis.py b/tests/test_feature_analysis.py new file mode 100644 index 0000000..60e4231 --- /dev/null +++ b/tests/test_feature_analysis.py @@ -0,0 +1,122 @@ +import json +import unittest +from unittest.mock import patch + +import numpy as np +import torch + +from test_model_loading import load_nodes +from feature_analysis import CHART_TYPES, analyze_features, prepare_features + + +class FeatureAnalysisTests(unittest.TestCase): + @classmethod + def setUpClass(cls): + torch.set_num_threads(2) + cls.vectors=torch.tensor([[1.,0.,0.],[.9,.1,0.],[0.,1.,0.],[0.,.9,.1],[0.,0.,1.],[.1,0.,.9]]) + cls.nodes=load_nodes() + + def test_every_chart_returns_finite_image_and_original_distance_matrix(self): + expected=1-torch.nn.functional.normalize(self.vectors,dim=1)@torch.nn.functional.normalize(self.vectors,dim=1).T + expected.fill_diagonal_(0) + for chart in CHART_TYPES: + with self.subTest(chart=chart): + features=self.vectors + options={} + if chart=='patch_energy': + features=self.vectors[:,None,:].repeat(1,4,1) + options['tensor_layout']='tokens' + image,text,distances=analyze_features(features,chart_type=chart,width=800,height=600,**options) + report=json.loads(text) + self.assertEqual(tuple(image.shape),(1,600,800,3)) + self.assertTrue(torch.isfinite(image).all()) + self.assertGreater(float(image.max()-image.min()),0.5) + torch.testing.assert_close(distances,expected,atol=2e-7,rtol=1e-6) + self.assertEqual(report['chart_type'],chart) + self.assertEqual(len(report['nearest_neighbors']),6) + for index,row in enumerate(report['nearest_neighbors']): + self.assertNotIn(index,[item['index'] for item in row]) + + def test_graph_edges_are_knn_union_and_use_original_distance(self): + _,text,distances=analyze_features(self.vectors,top_k=1,width=800,height=600) + report=json.loads(text) + expected={tuple(sorted((i,row[0]['index']))) for i,row in enumerate(report['nearest_neighbors'])} + actual={(edge['source'],edge['target']) for edge in report['edges']} + self.assertEqual(actual,expected) + for edge in report['edges']: + self.assertAlmostEqual(edge['distance'],float(distances[edge['source'],edge['target']]),places=6) + + def test_distances_and_normalization_match_known_values(self): + features=torch.tensor([[3.,4.],[0.,5.]]) + for metric,value in (('euclidean',np.sqrt(10)),('manhattan',4.0),('cosine',0.2)): + _,_,dist=analyze_features(features,chart_type='distance_heatmap',metric=metric,normalize=False,width=800,height=600) + self.assertAlmostEqual(float(dist[0,1]),value,places=6) + _,text,dist=analyze_features(features,chart_type='distance_heatmap',metric='euclidean',normalize=True,width=800,height=600) + self.assertAlmostEqual(float(dist[0,1]),np.sqrt(.4),places=6) + self.assertEqual(json.loads(text)['normalization'],'row L2') + + def test_layout_reductions_preserve_image_batch(self): + tokens=torch.arange(2*3*4,dtype=torch.float32).reshape(2,3,4) + vector,patches,_,_=prepare_features(tokens,'tokens') + np.testing.assert_allclose(vector,tokens.mean(1).numpy()) + self.assertEqual(patches.shape,(2,3,4)) + spatial=torch.arange(2*4*2*3,dtype=torch.float32).reshape(2,4,2,3) + vector,patches,grid,_=prepare_features(spatial,'spatial') + np.testing.assert_allclose(vector,spatial.mean((2,3)).numpy()) + self.assertEqual(grid,(2,3)) + layers=torch.stack((tokens,tokens+10),dim=1) + vector,_,_,_=prepare_features(layers,'layer_tokens',layer_index=-1) + np.testing.assert_allclose(vector,(tokens+10).mean(1).numpy()) + vector,_,_,_=prepare_features(layers,'layer_tokens',layer_pooling='mean') + np.testing.assert_allclose(vector,(tokens+5).mean(1).numpy()) + vector,_,_,_=prepare_features(layers,'flatten') + self.assertEqual(vector.shape,(2,24)) + with self.assertRaisesRegex(ValueError,'ambiguous'): + prepare_features(layers) + + def test_bad_inputs_and_undefined_cosine_fail_explicitly(self): + for features in (torch.tensor([[0.,0.],[1.,0.]]),torch.tensor([[float('nan'),1.]])): + with self.assertRaises(ValueError): + analyze_features(features,width=800,height=600) + with self.assertRaisesRegex(ValueError,'labels'): + analyze_features(self.vectors,labels='one',width=800,height=600) + with self.assertRaisesRegex(ValueError,'matching features'): + analyze_features(self.vectors,images=torch.ones(1,32,32,3),width=800,height=600) + with self.assertRaisesRegex(ValueError,'patch_energy requires'): + analyze_features(self.vectors,chart_type='patch_energy',width=800,height=600) + + def test_degenerate_inputs_do_not_fabricate_silhouette_or_pca(self): + identical=torch.ones(4,3) + _,text,dist=analyze_features(identical,chart_type='silhouette',width=800,height=600) + report=json.loads(text) + self.assertIsNone(report['silhouette_scores']) + self.assertEqual(report['pca_explained_variance_ratio'],[0.,0.,0.]) + self.assertTrue(report['warnings']) + torch.testing.assert_close(dist,torch.zeros(4,4)) + _,text,_=analyze_features(torch.zeros(3,2),metric='euclidean',chart_type='dimension_correlation',width=800,height=600) + self.assertEqual(json.loads(text)['dimension_correlations'],[[None,None],[None,None]]) + + def test_feature_node_never_invokes_model_and_image_node_invokes_extraction_once(self): + node=self.nodes.KaloscopeFeatureAnalysisNode() + self.assertNotIn('model',node.INPUT_TYPES()['required']) + with patch.object(self.nodes.KaloscopeExtractFeaturesNode,'extract',side_effect=AssertionError('unexpected inference')): + image,_,matrix=node.analyze(self.vectors,chart_type='relationship_graph',width=800,height=600) + self.assertEqual(matrix.shape,(6,6)) + with patch.object(self.nodes.KaloscopeExtractFeaturesNode,'extract',return_value=(self.vectors,)) as extraction: + outputs=self.nodes.KaloscopeImageAnalysisNode().analyze(torch.ones(6,32,32,3),{},width=800,height=600) + extraction.assert_called_once() + torch.testing.assert_close(outputs[2],self.vectors) + self.assertEqual(tuple(outputs[0].shape),(1,600,800,3)) + + def test_clustering_options_and_seed_are_reproducible(self): + for method in ('kmeans','agglomerative','dbscan','none'): + _,text,_=analyze_features(self.vectors,chart_type='cluster_sizes',cluster_method=method,width=800,height=600) + self.assertEqual(len(json.loads(text)['cluster_labels']),6) + _,first,_=analyze_features(self.vectors,chart_type='pca_scatter',seed=7,width=800,height=600) + _,second,_=analyze_features(self.vectors,chart_type='pca_scatter',seed=7,width=800,height=600) + self.assertEqual(json.loads(first)['cluster_labels'],json.loads(second)['cluster_labels']) + self.assertEqual(json.loads(first)['projection'],json.loads(second)['projection']) + + +if __name__=='__main__': + unittest.main() diff --git a/tests/test_model_loading.py b/tests/test_model_loading.py new file mode 100644 index 0000000..d3422bf --- /dev/null +++ b/tests/test_model_loading.py @@ -0,0 +1,394 @@ +"""Functional tests with real, small DINOv3 networks and local checkpoints.""" +import base64 +import importlib.util +import io +import json +from pathlib import Path +import sys +import tempfile +import types +import unittest +from unittest.mock import patch + +import torch +from PIL import Image + +ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(ROOT)) + +from model_loading import load_model_bundle, find_checkpoint, model_folders +from inference_artist import classify_image, resolved_mode +from kaloscope_dinov3.models.vision_transformer import DinoVisionTransformer + + +class ReferenceInferenceModel(torch.nn.Module): + """Independent forward-only fixture for comparing loaded checkpoint outputs.""" + def __init__(self, config, num_classes): + super().__init__() + self.backbone = DinoVisionTransformer(**config["kwargs"]) + self.backbone.init_weights() + self.feature_dim = 2 * self.backbone.embed_dim + self.head = torch.nn.Linear(self.feature_dim, num_classes) + self.projector = torch.nn.Sequential( + torch.nn.Linear(self.feature_dim, self.feature_dim), torch.nn.GELU(), + torch.nn.Linear(self.feature_dim, config["projection_dim"]), + ) + + def forward(self, images, projections=False): + tokens = self.backbone.forward_features(images) + features = torch.cat((tokens["x_norm_clstoken"], tokens["x_norm_patchtokens"].mean(1)), dim=1) + result = {"features": features, "logits": self.head(features)} + if projections: + result["projections"] = self.projector(features) + return result + + +def load_nodes(): + fake = types.ModuleType("folder_paths") + fake.models_dir = str(ROOT / "models") + spec = importlib.util.spec_from_file_location("kaloscope_test_nodes", ROOT / "__init__.py") + module = importlib.util.module_from_spec(spec) + original = sys.modules.get("folder_paths") + sys.modules["folder_paths"] = fake + try: + spec.loader.exec_module(module) + finally: + if original is None: + del sys.modules["folder_paths"] + else: + sys.modules["folder_paths"] = original + return module + + +class ModelLoadingTests(unittest.TestCase): + @classmethod + def setUpClass(cls): + torch.set_num_threads(2) + cls.options = {"name": "custom_vit", "pooling": "cls_mean", "projection_dim": 8, + "kwargs": {"embed_dim": 24, "depth": 2, "num_heads": 3, "patch_size": 16, + "n_storage_tokens": 2, + "pos_embed_rope_dtype": "fp32"}} + torch.manual_seed(1) + cls.original = ReferenceInferenceModel(cls.options, num_classes=3).eval() + cls.nodes = load_nodes() + + def setUp(self): + self.temp = tempfile.TemporaryDirectory() + self.directory = Path(self.temp.name) + self.config = {"model": "custom_vit", "input_size": 32} + self.payload = {"model": self.original.state_dict(), "model_config": self.options, + "classes": ["a", "b", "c"]} + self.save() + + def tearDown(self): + self.temp.cleanup() + + def save(self): + (self.directory / "config.json").write_text(json.dumps(self.config), encoding="utf-8") + torch.save(self.payload, self.directory / "best.pt") + + def bundle(self): + return load_model_bundle(self.directory, device="cpu") + + def tensor(self, bundle): + return bundle["transform"](Image.new("RGB", (48, 32), (70, 130, 190))).unsqueeze(0) + + def test_head_and_features_match_reference_model(self): + bundle = self.bundle() + batch = self.tensor(bundle) + with torch.inference_mode(): + expected = self.original(batch) + torch.testing.assert_close(bundle["model"](batch), expected["logits"]) + torch.testing.assert_close(bundle["model"](batch, return_features=True), expected["features"]) + self.assertTrue(bundle["has_classifier"]) + self.assertEqual(bundle["feature_dim"], 48) + results = classify_image(bundle["model"], batch, "cpu", bundle["class_mapping"], top_k=20) + self.assertEqual(len(results), 3) + self.assertEqual({row["class_name"] for row in results}, {"a", "b", "c"}) + self.assertAlmostEqual(sum(row["probability"] for row in results), 1.0, places=6) + + def test_no_head_is_feature_only_even_with_metadata_classes(self): + self.payload["model"] = {k: v for k, v in self.payload["model"].items() if not k.startswith("head.")} + self.save() + bundle = self.bundle() + self.assertFalse(bundle["has_classifier"]) + self.assertEqual(resolved_mode(bundle["model"], "auto"), "cluster") + with self.assertRaisesRegex(ValueError, "no classification head"): + classify_image(bundle["model"], self.tensor(bundle), "cpu") + with torch.inference_mode(): + torch.testing.assert_close(bundle["model"](self.tensor(bundle), return_features=True), + self.original(self.tensor(bundle))["features"]) + + def test_temporal_projector_is_loaded_and_used(self): + self.payload["model"] = {k: v for k, v in self.payload["model"].items() if not k.startswith("head.")} + self.payload["model"].update(log_temperature=torch.tensor(1.0), bias=torch.tensor(-1.0)) + self.save() + bundle = self.bundle() + self.assertEqual(bundle["feature_source"], "projector") + self.assertEqual(bundle["feature_dim"], 8) + with torch.inference_mode(): + expected = self.original(self.tensor(bundle), projections=True)["projections"] + torch.testing.assert_close(bundle["model"](self.tensor(bundle), return_features=True), expected) + + def test_raw_backbone_and_official_linear_head(self): + self.payload = dict(self.original.backbone.state_dict()) + self.payload.update({"linear_head.weight": self.original.head.weight, + "linear_head.bias": self.original.head.bias}) + self.config["model"] = {"name": "custom_vit", "kwargs": self.options["kwargs"]} + self.save() + bundle = self.bundle() + self.assertEqual(bundle["model"].pooling, "cls_mean") + with torch.inference_mode(): + torch.testing.assert_close(bundle["model"](self.tensor(bundle)), self.original(self.tensor(bundle))["logits"]) + self.assertTrue(classify_image(bundle["model"], self.tensor(bundle), "cpu")[0]["class_name"].startswith("Class ")) + + def test_raw_headless_safetensors(self): + from safetensors.torch import save_file + self.config["model"] = {"name": "custom_vit", "kwargs": self.options["kwargs"]} + self.config["checkpoint"] = "backbone.safetensors" + self.save() + save_file(self.original.backbone.state_dict(), self.directory / "backbone.safetensors") + bundle = self.bundle() + self.assertFalse(bundle["has_classifier"]) + self.assertEqual(bundle["feature_dim"], 24) + with torch.inference_mode(): + expected = self.original.backbone.forward_features(self.tensor(bundle))["x_norm_clstoken"] + torch.testing.assert_close(bundle["model"](self.tensor(bundle), return_features=True), expected) + + def test_wrapped_state_dict_prefixes(self): + self.payload = {"state_dict": {"module._orig_mod." + k: v for k, v in self.payload["model"].items()}, + "model_config": self.options} + self.save() + self.assertTrue(self.bundle()["has_classifier"]) + + def test_checkpoint_transform_path_works_without_finetune_module(self): + from kaloscope_dinov3.preprocessing import lvd_transform + image = Image.new("RGB", (48, 32), "blue") + for prefix in ("dinov3.finetune.data.", "kaloscope_dinov3.finetune.data."): + self.config["data"] = {"custom_transform": prefix + "lvd_transform"} + self.save() + torch.testing.assert_close(self.bundle()["transform"](image), lvd_transform(32)(image)) + + def test_all_final_feature_outputs_match_backbone(self): + bundle = self.bundle() + model = bundle['model'] + batch = self.tensor(bundle) + with torch.inference_mode(): + raw = self.original.backbone.forward_features(batch) + cls, patches, storage = raw['x_norm_clstoken'], raw['x_norm_patchtokens'], raw['x_storage_tokens'] + pooled = torch.cat((cls, patches.mean(1)), dim=1) + expected = { + 'default': pooled, 'backbone': pooled, 'cls': cls, 'mean': patches.mean(1), + 'cls_mean': pooled, 'patch_tokens': patches, 'storage_tokens': storage, + 'all_tokens': torch.cat((cls.unsqueeze(1), storage, patches), dim=1), + 'prenorm': raw['x_prenorm'], 'projector': self.original.projector(pooled), + 'patch_map': patches.transpose(1, 2).reshape(1, 24, 2, 2), + } + for kind, tensor in expected.items(): + with self.subTest(output=kind): + torch.testing.assert_close(model.extract_tensor(batch, kind), tensor) + # Loading backbone as default must still retain trained projector for selection. + self.assertEqual(bundle['feature_source'], 'backbone') + self.assertIsNotNone(model.projector) + + def test_intermediate_features_keep_layer_order_and_shape(self): + model = self.bundle()['model'] + batch = self.tensor(self.bundle()) + with torch.inference_mode(): + native = self.original.backbone.get_intermediate_layers( + batch, n=[0, 1], return_class_token=True, return_extra_tokens=True) + for kind in ('cls', 'mean', 'cls_mean', 'patch_tokens', 'patch_map', 'storage_tokens', 'all_tokens'): + expected = [] + for patches, cls, storage in reversed(native): + if kind == 'cls': + tensor = cls + elif kind == 'mean': + tensor = patches.mean(1) + elif kind == 'cls_mean': + tensor = torch.cat((cls, patches.mean(1)), dim=1) + elif kind == 'patch_tokens': + tensor = patches + elif kind == 'patch_map': + tensor = patches.transpose(1, 2).reshape(1, 24, 2, 2) + elif kind == 'storage_tokens': + tensor = storage + else: + tensor = torch.cat((cls.unsqueeze(1), storage, patches), dim=1) + expected.append(tensor) + with self.subTest(output=kind): + actual = model.extract_tensor(batch, 'intermediate_' + kind, '-1,0') + torch.testing.assert_close(actual, torch.stack(expected, dim=1)) + unnorm = self.original.backbone.get_intermediate_layers( + batch, n=[0, 1], norm=False, return_class_token=True, return_extra_tokens=True) + expected = torch.stack([torch.cat((cls.unsqueeze(1), storage, patches), dim=1) + for patches, cls, storage in unnorm], dim=1) + torch.testing.assert_close(model.extract_tensor(batch, 'intermediate_prenorm', '0,1'), expected) + torch.testing.assert_close(model.extract_tensor(batch, 'intermediate_all_tokens', '0,1', False), expected) + self.assertEqual(tuple(model.extract_tensor(batch, 'intermediate_patch_tokens').shape), (1, 1, 4, 24)) + + def test_feature_selection_does_not_change_classifier(self): + bundle = self.bundle() + model = bundle['model'] + batch = self.tensor(bundle) + with torch.inference_mode(): + before = model(batch) + for kind in ('mean', 'projector', 'patch_tokens', 'cls'): + model.extract_tensor(batch, kind) + torch.testing.assert_close(model(batch), before) + self.assertEqual(model.pooling, 'cls_mean') + self.assertEqual(model.feature_source, 'backbone') + + def test_missing_projector_and_invalid_layers_fail_clearly(self): + self.payload['model'] = {k: v for k, v in self.payload['model'].items() if not k.startswith('projector.')} + self.save() + bundle = self.bundle() + batch = self.tensor(bundle) + with self.assertRaisesRegex(ValueError, 'no supported projector'): + bundle['model'].extract_tensor(batch, 'projector') + for layers in ('2', '-3', '0,0', 'one', ''): + with self.subTest(layers=layers), self.assertRaises(ValueError): + bundle['model'].extract_tensor(batch, 'intermediate_cls', layers) + + def test_feature_node_returns_selected_tensor_for_image_batch(self): + bundle = self.bundle() + image = torch.full((2, 32, 48, 3), 0.5) + node = self.nodes.KaloscopeExtractFeaturesNode() + for kind, shape in { + 'cls': (2, 24), 'mean': (2, 24), 'cls_mean': (2, 48), 'projector': (2, 8), + 'patch_tokens': (2, 4, 24), 'patch_map': (2, 24, 2, 2), + 'storage_tokens': (2, 2, 24), 'all_tokens': (2, 7, 24), 'prenorm': (2, 7, 24), + 'intermediate_patch_tokens': (2, 2, 4, 24), + 'intermediate_patch_map': (2, 2, 24, 2, 2), + }.items(): + with self.subTest(output=kind): + result = node.extract(image, bundle, kind, '0,1')[0] + self.assertEqual(tuple(result.shape), shape) + self.assertEqual(result.device.type, 'cpu') + self.assertFalse(result.requires_grad) + self.assertTrue(torch.isfinite(result).all()) + + def test_convnext_feature_outputs_and_variable_stages(self): + from kaloscope_dinov3.models.convnext import ConvNeXt + from model_loading import DinoInferenceModel + backbone = ConvNeXt(depths=[1, 1, 1, 1], dims=[8, 16, 24, 32]).eval() + model = DinoInferenceModel(backbone, 'cls').eval() + images = torch.rand(2, 3, 64, 96) + with torch.inference_mode(): + raw = backbone.forward_features(images) + torch.testing.assert_close(model.extract_tensor(images, 'patch_tokens'), raw['x_norm_patchtokens']) + patch_map = model.extract_tensor(images, 'patch_map') + self.assertEqual(tuple(patch_map.shape), (2, 32, 2, 3)) + torch.testing.assert_close(patch_map.flatten(2).transpose(1, 2), raw['x_norm_patchtokens']) + intermediate = model.extract_tensor(images, 'intermediate_patch_map', '1') + self.assertEqual(tuple(intermediate.shape), (2, 1, 16, 8, 12)) + torch.testing.assert_close(model.extract_tensor(images, 'intermediate_prenorm', '-1')[:, 0], raw['x_prenorm']) + with self.assertRaisesRegex(ValueError, 'no storage/register tokens'): + model.extract_tensor(images, 'storage_tokens') + with self.assertRaisesRegex(ValueError, 'different tensor shapes'): + model.extract_tensor(images, 'intermediate_patch_tokens', '0,1') + + def test_lsnet_feature_node_preserves_default_and_rejects_dino_outputs(self): + encoder = torch.nn.Identity() + encoder.forward = lambda batch, return_features: batch.mean((2, 3)) + from kaloscope_dinov3.preprocessing import image_transform + bundle = {'model': encoder, 'transform': image_transform(32), 'device': 'cpu'} + node = self.nodes.KaloscopeExtractFeaturesNode() + image = torch.ones(1, 32, 32, 3) + self.assertEqual(tuple(node.extract(image, bundle)[0].shape), (1, 3)) + with self.assertRaisesRegex(ValueError, 'requires a DINOv3'): + node.extract(image, bundle, 'patch_tokens') + + def test_bad_architecture_is_not_silently_ignored(self): + self.config["model"] = "dinov3_typo" + self.save() + with self.assertRaises(ValueError): + self.bundle() + + def test_incomplete_backbone_fails(self): + del self.payload["model"]["backbone.cls_token"] + self.save() + with self.assertRaisesRegex(RuntimeError, "cls_token"): + self.bundle() + + def test_class_mapping_must_match_head(self): + (self.directory / "class_mapping.csv").write_text("class_id,class_name\n0,one\n2,two\n", encoding="utf-8") + with self.assertRaisesRegex(ValueError, "mapping IDs"): + self.bundle() + + def test_model_discovery_only_uses_kaloscope(self): + for folder in ("lsnet/shared", "kaloscope/shared", "lsnet/other"): + (self.directory / folder).mkdir(parents=True) + folders = model_folders(self.directory) + self.assertEqual(folders, {"shared": self.directory / "kaloscope/shared"}) + + def test_architecture_must_be_explicit(self): + del self.config['model'] + self.save() + with self.assertRaisesRegex(ValueError, 'model architecture in config.json'): + self.bundle() + bundle = load_model_bundle(self.directory, device='cpu', model_name='custom_vit') + self.assertEqual(bundle['model_type'], 'custom_vit') + + def test_ambiguous_checkpoint_requires_selection(self): + (self.directory / "best.pt").rename(self.directory / "one.pt") + torch.save(self.payload, self.directory / "two.pth") + with self.assertRaisesRegex(ValueError, "Expected one checkpoint"): + find_checkpoint(self.directory) + self.config["checkpoint"] = "two.pth" + self.save() + self.assertEqual(find_checkpoint(self.directory).name, "two.pth") + + def test_kaloscope_nodes_and_model_sockets(self): + bundle = self.bundle() + nodes = self.nodes + self.assertEqual(len(nodes.NODE_CLASS_MAPPINGS), 10) + self.assertEqual(set(nodes.NODE_CLASS_MAPPINGS), set(nodes.NODE_DISPLAY_NAME_MAPPINGS)) + self.assertEqual(nodes.KaloscopeModelLoader.RETURN_TYPES, ("KALOSCOPE_MODEL",)) + for name, node in nodes.NODE_CLASS_MAPPINGS.items(): + schema = node.INPUT_TYPES() + self.assertTrue(name.startswith('Kaloscope')) + if "model" in schema.get("required", {}): + self.assertEqual(schema["required"]["model"][0], 'KALOSCOPE_MODEL') + self.assertTrue(node.__name__.startswith("Kaloscope")) + self.assertIn(node.CATEGORY, ("Kaloscope", "Kaloscope/Analysis")) + with patch.object(nodes, "model_folders", return_value={"sample": self.directory}): + loaded = nodes.KaloscopeModelLoader().load("sample", "cpu")[0] + self.assertTrue(loaded["has_classifier"]) + image = torch.full((2, 32, 48, 3), 0.5) + features = nodes.KaloscopeExtractFeaturesNode().extract(image, bundle)[0] + self.assertEqual(tuple(features.shape), (2, 48)) + tags, predictions = nodes.KaloscopeArtistInferenceNode().process(image, bundle, 5, 0.0) + self.assertEqual(set(tags.split(",")), {"a", "b", "c"}) + self.assertEqual(len(json.loads(predictions)), 3) + result, visualization = nodes.KaloscopeClusteringNode().cluster( + "kmeans", 2, 0.5, 2, True, "pca", 5, group_1=torch.randn(4, 48)) + self.assertEqual(len(json.loads(result)["labels"]), 4) + self.assertEqual(visualization.shape[-1], 3) + + def test_backend_auto_features_and_api(self): + from backend_lsnet.inference import process_image_from_pil + from backend_lsnet.api import Api + from fastapi import FastAPI + from fastapi.testclient import TestClient + self.payload["model"] = {k: v for k, v in self.payload["model"].items() if not k.startswith("head.")} + self.save() + image = Image.new("RGB", (48, 32), "blue") + result = process_image_from_pil(image, checkpoint=str(self.directory / "best.pt"), device="cpu") + self.assertEqual(len(result["features"]), 48) + buffer = io.BytesIO() + image.save(buffer, format="PNG") + app = FastAPI() + api = Api(app) + with patch("backend_lsnet.api.get_available_checkpoints", return_value=["best.pt"]), \ + patch("backend_lsnet.api.get_checkpoint_path", return_value=str(self.directory / "best.pt")), \ + patch("backend_lsnet.api.get_class_csv", return_value=None), TestClient(app) as client: + response = client.post('/kaloscope/v1/infer', json={"input_image": base64.b64encode(buffer.getvalue()).decode(), "device": "cpu"}) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(len(response.json()["results"]["features"]), 48) + self.assertTrue(all(route.path.startswith('/kaloscope/v1/') for route in app.routes + if hasattr(route, 'endpoint') and route.endpoint.__module__ == 'backend_lsnet.api')) + api.executor.shutdown() + + +if __name__ == "__main__": + unittest.main() diff --git a/单独启动.bat b/单独启动.bat index 1938e1a..d95ce83 100644 --- a/单独启动.bat +++ b/单独启动.bat @@ -1,2 +1,4 @@ +@echo off +cd /d "%~dp0" python -m scripts.app -pause \ No newline at end of file +pause