修改配置和路径。

This commit is contained in:
2025-10-23 11:08:46 +08:00
parent 5db96600fe
commit 1c8ceef5b0
5 changed files with 29 additions and 21 deletions
+2 -2
View File
@@ -65,8 +65,8 @@ class EmbeddingFoodClassifierApp:
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
# model_path = "../model/embedding_20251011_133653/best_embedding_model.pth"
# model_path = os.path.join(BASE_DIR, "../model/embedding_20251011_133653/best_embedding_model.pth")
model_path = os.path.join(BASE_DIR, "../model/WholeIngredientRecognition/embedding_20251017_145836/best_embedding_model.pth")
# model_path = os.path.join(BASE_DIR, "../model/DishClassification/embedding_20251011_133653/best_embedding_model.pth")
model_path = os.path.join(BASE_DIR, "../model/WholeIngredientRecognition/embedding_20251021_085915/best_embedding_model.pth")
# model_path = os.path.join(BASE_DIR, "../model/DishClassification/embedding_20251022_093635/best_embedding_model.pth")
# FAISS索引目录
# index_dir = "../faiss_vector_db/faiss_index"
+6 -6
View File
@@ -480,12 +480,12 @@ def main():
"""主函数"""
# 配置参数
# 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"
# MODEL_PATH = "../model/WholeIngredientRecognition/embedding_20251021_085915/best_embedding_model.pth"
MODEL_PATH = "../model/DishClassification/embedding_20251022_093635/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
+4 -4
View File
@@ -162,10 +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("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("--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为全量")
+2 -2
View File
@@ -18,7 +18,7 @@ def main():
# base_model = create_resnet50_embedding(embedding_dim=512, pretrained=True)
base_model = create_mobile_resnet50_embedding(embedding_dim=512, pretrained=True)
# model_path = "../model/embedding_20250930_102826/best_embedding_model.pth"
model_path = "../model/embedding_20251011_133653/best_embedding_model.pth"
model_path = "../model/WholeIngredientRecognition/embedding_20251021_085915/best_embedding_model.pth"
if not os.path.exists(model_path):
print(f"错误:模型文件不存在 {model_path}")
@@ -86,7 +86,7 @@ def main():
traced_model = torch.jit.trace(mobile_wrapper, single_input)
# 保存模型
output_path = "../model/embedding_20251011_133653/best_embedding_model_mobile.pt"
output_path = "../model/WholeIngredientRecognition/embedding_20251021_085915/best_embedding_model_mobile.pt"
traced_model.save(output_path)
print(f"✓ TorchScript模型保存成功: {output_path}")
+15 -7
View File
@@ -48,8 +48,10 @@ TASKS = {
embedding_dim=512,
batch_size=16,
lr=1e-3,
triplet_margin=0.3,
center_loss_weight=0.1,
# triplet_margin=0.3,
triplet_margin=0.5,
# center_loss_weight=0.1,
center_loss_weight=0.5,
aug_strength="medium",
),
"whole_ingredient": TaskConfig(
@@ -59,7 +61,8 @@ TASKS = {
embedding_dim=512,
batch_size=32,
lr=8e-4,
triplet_margin=0.35,
# triplet_margin=0.35,
triplet_margin=0.5,
center_loss_weight=0.1,
aug_strength="medium",
),
@@ -385,7 +388,7 @@ class EarlyStopping:
def train_epoch(model, train_loader, triplet_criterion, center_criterion,
optimizer, center_optimizer, device, epoch):
optimizer, center_optimizer, device, epoch,CENTER_LOSS_WEIGHT):
"""
训练一个epoch
@@ -430,7 +433,11 @@ def train_epoch(model, train_loader, triplet_criterion, center_criterion,
center_loss = center_criterion(anchor_emb, labels)
# 总损失
loss = triplet_loss + settings.CENTER_LOSS_WEIGHT * center_loss # 可配置的中心损失权重
# loss = triplet_loss + settings.CENTER_LOSS_WEIGHT * center_loss # 可配置的中心损失权重
loss = triplet_loss + CENTER_LOSS_WEIGHT * center_loss # 可配置的中心损失权重
# print('triplet_loss',triplet_loss)
# print('center_loss',CENTER_LOSS_WEIGHT * center_loss)
# print('CENTER_LOSS_WEIGHT',CENTER_LOSS_WEIGHT)
# 反向传播
optimizer.zero_grad()
@@ -709,7 +716,7 @@ def main(task_key: str = "dish"):
# 训练
train_loss, train_triplet_loss, train_center_loss = train_epoch(
model, train_loader, triplet_criterion, center_criterion,
optimizer, center_optimizer, device, epoch
optimizer, center_optimizer, device, epoch,CENTER_LOSS_WEIGHT
)
# 验证
@@ -811,4 +818,5 @@ if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("task", choices=list(TASKS.keys()), nargs="?", default="dish")
args = parser.parse_args()
main("whole_ingredient")
# main("whole_ingredient")
main("dish")