mine_layer3_step1_extract_v0904.py 3.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102
  1. # -*- coding: utf-8 -*-
  2. """mine_layer3_step1_extract.py - 用老下半区模型对全库下半区图提特征(挖矿原料)。
  3. 输入: lower_all_data_v0904/images/*.jpg
  4. 输出: lower_all_feats_v0904.npy (n,1024) + lower_all_ids_v0904.json
  5. """
  6. import os
  7. os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
  8. import json
  9. import time
  10. import cv2
  11. import numpy as np
  12. import torch
  13. import torch.nn as nn
  14. import torch.nn.functional as F
  15. from torch.utils.data import DataLoader, Dataset
  16. from transformers import Dinov2Model
  17. WZJ = "/home/user/顾工交接/wzj"
  18. IMG_DIR = os.path.join(WZJ, "lower_all_data_v0904/images")
  19. MODEL_PATH = os.path.join(WZJ, "layer3_bg_model_output/best_layer3_bottom_model.pth")
  20. OUT_NPY = os.path.join(WZJ, "lower_all_feats_v0904.npy")
  21. OUT_IDS = os.path.join(WZJ, "lower_all_ids_v0904.json")
  22. HF_MODEL_ID = "facebook/dinov2-large"
  23. IMG_HEIGHT, IMG_WIDTH = 196, 392
  24. FREEZE_BLOCKS = 18
  25. BATCH = 128
  26. DEVICE = torch.device("cuda:0")
  27. class Dinov2Layer3Model(nn.Module):
  28. def __init__(self, freeze_blocks=FREEZE_BLOCKS):
  29. super().__init__()
  30. self.backbone = Dinov2Model.from_pretrained(HF_MODEL_ID)
  31. for p in self.backbone.parameters():
  32. p.requires_grad = False
  33. for i in range(freeze_blocks, len(self.backbone.encoder.layer)):
  34. for p in self.backbone.encoder.layer[i].parameters():
  35. p.requires_grad = True
  36. for p in self.backbone.layernorm.parameters():
  37. p.requires_grad = True
  38. def forward(self, x):
  39. out = self.backbone(x, interpolate_pos_encoding=True)
  40. return F.normalize(out.last_hidden_state[:, 0, :], p=2, dim=1)
  41. class ImgDataset(Dataset):
  42. def __init__(self, files):
  43. self.files = files
  44. def __len__(self):
  45. return len(self.files)
  46. def __getitem__(self, i):
  47. img = cv2.imread(os.path.join(IMG_DIR, self.files[i]))
  48. if img is None:
  49. img = np.zeros((IMG_HEIGHT, IMG_WIDTH, 3), dtype=np.uint8)
  50. else:
  51. img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
  52. img = cv2.resize(img, (IMG_WIDTH, IMG_HEIGHT), interpolation=cv2.INTER_CUBIC)
  53. img = (img.astype(np.float32) / 255.0 - [0.485, 0.456, 0.406]) / [0.229, 0.224, 0.225]
  54. return torch.from_numpy(img.astype(np.float32).transpose(2, 0, 1))
  55. def main():
  56. files = sorted(f for f in os.listdir(IMG_DIR) if f.endswith(".jpg"))
  57. print(f"[files] {len(files)}", flush=True)
  58. ids = [f[:-4] for f in files]
  59. print("[model] load old layer3 weights ...", flush=True)
  60. model = Dinov2Layer3Model()
  61. model.load_state_dict(torch.load(MODEL_PATH, map_location="cpu"))
  62. model.to(DEVICE).eval()
  63. loader = DataLoader(ImgDataset(files), batch_size=BATCH, shuffle=False,
  64. num_workers=16, pin_memory=True)
  65. feats = np.zeros((len(files), 1024), dtype=np.float32)
  66. t0 = time.time()
  67. i0 = 0
  68. with torch.no_grad():
  69. for bi, batch in enumerate(loader):
  70. batch = batch.to(DEVICE, non_blocking=True)
  71. with torch.amp.autocast("cuda"):
  72. f = model(batch)
  73. feats[i0:i0 + batch.shape[0]] = f.float().cpu().numpy()
  74. i0 += batch.shape[0]
  75. if bi % 50 == 0:
  76. el = (time.time() - t0) / 60
  77. print(f" {i0}/{len(files)} elapsed={el:.1f}m", flush=True)
  78. np.save(OUT_NPY, feats)
  79. with open(OUT_IDS, "w", encoding="utf-8") as f:
  80. json.dump(ids, f, ensure_ascii=False)
  81. print(f"[DONE] {i0} feats -> {OUT_NPY} elapsed={(time.time()-t0)/60:.1f}m", flush=True)
  82. if __name__ == "__main__":
  83. main()