修改配置和路径。
This commit is contained in:
@@ -65,8 +65,8 @@ class EmbeddingFoodClassifierApp:
|
|||||||
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||||
# model_path = "../model/embedding_20251011_133653/best_embedding_model.pth"
|
# 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/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/WholeIngredientRecognition/embedding_20251021_085915/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/DishClassification/embedding_20251022_093635/best_embedding_model.pth")
|
||||||
|
|
||||||
# FAISS索引目录
|
# FAISS索引目录
|
||||||
# index_dir = "../faiss_vector_db/faiss_index"
|
# index_dir = "../faiss_vector_db/faiss_index"
|
||||||
|
|||||||
@@ -480,12 +480,12 @@ def main():
|
|||||||
"""主函数"""
|
"""主函数"""
|
||||||
# 配置参数
|
# 配置参数
|
||||||
# MODEL_PATH = "../model/embedding_20251011_133653/best_embedding_model.pth"
|
# MODEL_PATH = "../model/embedding_20251011_133653/best_embedding_model.pth"
|
||||||
MODEL_PATH = "../model/WholeIngredientRecognition/embedding_20251017_145836/best_embedding_model.pth"
|
# MODEL_PATH = "../model/WholeIngredientRecognition/embedding_20251021_085915/best_embedding_model.pth"
|
||||||
# MODEL_PATH = "../model/DishClassification/embedding_20251011_133653/best_embedding_model.pth"
|
MODEL_PATH = "../model/DishClassification/embedding_20251022_093635/best_embedding_model.pth"
|
||||||
TRAIN_DIR = "../dataset/WholeIngredientRecognition/train"
|
# TRAIN_DIR = "../dataset/WholeIngredientRecognition/train"
|
||||||
# TRAIN_DIR = "../dataset/DishClassification/train"
|
TRAIN_DIR = "../dataset/DishClassification/train"
|
||||||
OUTPUT_DIR = "WholeIngredientRecognition/faiss_index"
|
# OUTPUT_DIR = "WholeIngredientRecognition/faiss_index"
|
||||||
# OUTPUT_DIR = "DishClassification/faiss_index"
|
OUTPUT_DIR = "DishClassification/faiss_index"
|
||||||
BATCH_SIZE = 16
|
BATCH_SIZE = 16
|
||||||
INDEX_TYPE = 'flat' # 'flat', 'ivf', 'hnsw'
|
INDEX_TYPE = 'flat' # 'flat', 'ivf', 'hnsw'
|
||||||
EMBEDDING_DIM = 512
|
EMBEDDING_DIM = 512
|
||||||
|
|||||||
@@ -162,10 +162,10 @@ def plot_2d(Z: np.ndarray, y: np.ndarray, title: str, out_path: Optional[str] =
|
|||||||
|
|
||||||
def main():
|
def main():
|
||||||
parser = argparse.ArgumentParser(description="可视化高维 embedding 并进行坍塌诊断")
|
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("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("--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("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("--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("--method", type=str, default="pca", choices=["pca", "tsne", "umap"], help="降维方法")
|
||||||
parser.add_argument("--seed", type=int, default=42)
|
parser.add_argument("--seed", type=int, default=42)
|
||||||
parser.add_argument("--max_points", type=int, default=None, help="抽样上限,避免t-SNE/UMAP过慢;None为全量")
|
parser.add_argument("--max_points", type=int, default=None, help="抽样上限,避免t-SNE/UMAP过慢;None为全量")
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ def main():
|
|||||||
# base_model = create_resnet50_embedding(embedding_dim=512, pretrained=True)
|
# base_model = create_resnet50_embedding(embedding_dim=512, pretrained=True)
|
||||||
base_model = create_mobile_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_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):
|
if not os.path.exists(model_path):
|
||||||
print(f"错误:模型文件不存在 {model_path}")
|
print(f"错误:模型文件不存在 {model_path}")
|
||||||
@@ -86,7 +86,7 @@ def main():
|
|||||||
traced_model = torch.jit.trace(mobile_wrapper, single_input)
|
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)
|
traced_model.save(output_path)
|
||||||
print(f"✓ TorchScript模型保存成功: {output_path}")
|
print(f"✓ TorchScript模型保存成功: {output_path}")
|
||||||
|
|
||||||
|
|||||||
@@ -48,8 +48,10 @@ TASKS = {
|
|||||||
embedding_dim=512,
|
embedding_dim=512,
|
||||||
batch_size=16,
|
batch_size=16,
|
||||||
lr=1e-3,
|
lr=1e-3,
|
||||||
triplet_margin=0.3,
|
# triplet_margin=0.3,
|
||||||
center_loss_weight=0.1,
|
triplet_margin=0.5,
|
||||||
|
# center_loss_weight=0.1,
|
||||||
|
center_loss_weight=0.5,
|
||||||
aug_strength="medium",
|
aug_strength="medium",
|
||||||
),
|
),
|
||||||
"whole_ingredient": TaskConfig(
|
"whole_ingredient": TaskConfig(
|
||||||
@@ -59,7 +61,8 @@ TASKS = {
|
|||||||
embedding_dim=512,
|
embedding_dim=512,
|
||||||
batch_size=32,
|
batch_size=32,
|
||||||
lr=8e-4,
|
lr=8e-4,
|
||||||
triplet_margin=0.35,
|
# triplet_margin=0.35,
|
||||||
|
triplet_margin=0.5,
|
||||||
center_loss_weight=0.1,
|
center_loss_weight=0.1,
|
||||||
aug_strength="medium",
|
aug_strength="medium",
|
||||||
),
|
),
|
||||||
@@ -385,7 +388,7 @@ class EarlyStopping:
|
|||||||
|
|
||||||
|
|
||||||
def train_epoch(model, train_loader, triplet_criterion, center_criterion,
|
def train_epoch(model, train_loader, triplet_criterion, center_criterion,
|
||||||
optimizer, center_optimizer, device, epoch):
|
optimizer, center_optimizer, device, epoch,CENTER_LOSS_WEIGHT):
|
||||||
"""
|
"""
|
||||||
训练一个epoch
|
训练一个epoch
|
||||||
|
|
||||||
@@ -430,7 +433,11 @@ def train_epoch(model, train_loader, triplet_criterion, center_criterion,
|
|||||||
center_loss = center_criterion(anchor_emb, labels)
|
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()
|
optimizer.zero_grad()
|
||||||
@@ -709,7 +716,7 @@ def main(task_key: str = "dish"):
|
|||||||
# 训练
|
# 训练
|
||||||
train_loss, train_triplet_loss, train_center_loss = train_epoch(
|
train_loss, train_triplet_loss, train_center_loss = train_epoch(
|
||||||
model, train_loader, triplet_criterion, center_criterion,
|
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 = argparse.ArgumentParser()
|
||||||
parser.add_argument("task", choices=list(TASKS.keys()), nargs="?", default="dish")
|
parser.add_argument("task", choices=list(TASKS.keys()), nargs="?", default="dish")
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
main("whole_ingredient")
|
# main("whole_ingredient")
|
||||||
|
main("dish")
|
||||||
Reference in New Issue
Block a user