增加一些设置。

This commit is contained in:
2025-10-17 18:24:04 +08:00
parent 25f8e0a267
commit 5db96600fe
10 changed files with 38 additions and 25 deletions
+1 -1
View File
@@ -111,7 +111,7 @@ builder = FAISSIndexBuilder(
# 构建完整索引
index = builder.build_complete_index(
train_dir="../dataset/train",
output_dir="faiss_index",
output_dir="DishClassification/faiss_index",
batch_size=16,
index_type='flat' # 'flat', 'ivf', 'hnsw'
)
+13 -7
View File
@@ -133,10 +133,12 @@ class FAISSIndexBuilder:
print(f"类别 '{class_name}': {len(image_files)} 张图片")
for image_file in image_files:
image_path = os.path.join(class_dir, image_file)
image_paths.append(image_path)
class_names.append(class_name)
labels.append(class_idx)
# 可以不添加那些增强的图片
if image_file.startswith("img"):
image_path = os.path.join(class_dir, image_file)
image_paths.append(image_path)
class_names.append(class_name)
labels.append(class_idx)
print(f"总计扫描到 {len(image_paths)} 张图片")
return image_paths, class_names, labels
@@ -477,9 +479,13 @@ class FAISSSearcher:
def main():
"""主函数"""
# 配置参数
MODEL_PATH = "../model/embedding_20251011_133653/best_embedding_model.pth"
TRAIN_DIR = "../dataset/train"
OUTPUT_DIR = "faiss_index"
# MODEL_PATH = "../model/embedding_20251011_133653/best_embedding_model.pth"
MODEL_PATH = "../model/WholeIngredientRecognition/embedding_20251017_145836/best_embedding_model.pth"
# MODEL_PATH = "../model/DishClassification/embedding_20251011_133653/best_embedding_model.pth"
TRAIN_DIR = "../dataset/WholeIngredientRecognition/train"
# TRAIN_DIR = "../dataset/DishClassification/train"
OUTPUT_DIR = "WholeIngredientRecognition/faiss_index"
# OUTPUT_DIR = "DishClassification/faiss_index"
BATCH_SIZE = 16
INDEX_TYPE = 'flat' # 'flat', 'ivf', 'hnsw'
EMBEDDING_DIM = 512
+2 -2
View File
@@ -87,7 +87,7 @@ def run_search_demo():
from faiss_vector_db.build_faiss_index import FAISSSearcher
# 配置参数
index_dir = "faiss_index"
index_dir = "DishClassification/faiss_index"
model_path = "../model/embedding_20250917_145342/best_embedding_model.pth"
# 检查索引是否存在
@@ -206,7 +206,7 @@ def main():
print("✓ 文件检查通过")
# 检查是否已有索引
index_exists = os.path.exists("faiss_index/faiss_index.bin")
index_exists = os.path.exists("DishClassification/faiss_index/faiss_index.bin")
if index_exists:
print("✓ 发现已存在的FAISS索引")
+4 -2
View File
@@ -162,8 +162,10 @@ def plot_2d(Z: np.ndarray, y: np.ndarray, title: str, out_path: Optional[str] =
def main():
parser = argparse.ArgumentParser(description="可视化高维 embedding 并进行坍塌诊断")
parser.add_argument("--embeddings", type=str, default=os.path.join( "faiss_index", "embeddings.json"), help="embeddings.json 路径")
parser.add_argument("--labels", type=str, default=os.path.join("faiss_index", "labels.json"), help="labels.json 路径")
# parser.add_argument("--embeddings", type=str, default=os.path.join("DishClassification/faiss_index", "embeddings.json"), help="embeddings.json 路径")
parser.add_argument("--embeddings", type=str, default=os.path.join("WholeIngredientRecognition/faiss_index", "embeddings.json"), help="embeddings.json 路径")
# parser.add_argument("--labels", type=str, default=os.path.join("DishClassification/faiss_index", "labels.json"), help="labels.json 路径")
parser.add_argument("--labels", type=str, default=os.path.join("WholeIngredientRecognition/faiss_index", "labels.json"), help="labels.json 路径")
parser.add_argument("--method", type=str, default="pca", choices=["pca", "tsne", "umap"], help="降维方法")
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--max_points", type=int, default=None, help="抽样上限,避免t-SNE/UMAP过慢;None为全量")