-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathmatch_embedding.py
More file actions
198 lines (165 loc) · 7.92 KB
/
Copy pathmatch_embedding.py
File metadata and controls
198 lines (165 loc) · 7.92 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
import json
import re
import os
import torch
import torch.nn.functional as F
from torch import Tensor
from transformers import AutoTokenizer, AutoModel
LRC_PATH = "./autodl-tmp/lyrics.lrc"
SCENES_DIR = "./autodl-tmp/scenes/"
BEST_MATCHES_PATH = "./autodl-tmp/best_matches.txt"
SCENE_EMBEDDINGS_PATH = "./autodl-tmp/scene_embeddings.pt"
TOP_K = 3 # Save top K best matches
MODEL_NAME = 'Qwen/Qwen3-Embedding-8B'
MAX_LENGTH = 256
BATCH_SIZE = 64 # Batch size for encoding to avoid OOM
TASK_INSTRUCTION = 'Retrieve relevant clips matching the scene.'
def last_token_pool(last_hidden_states: Tensor, attention_mask: Tensor) -> Tensor:
"""Extract embeddings from the last token position."""
left_padding = (attention_mask[:, -1].sum() == attention_mask.shape[0])
if left_padding:
return last_hidden_states[:, -1]
else:
sequence_lengths = attention_mask.sum(dim=1) - 1
batch_size = last_hidden_states.shape[0]
return last_hidden_states[torch.arange(batch_size, device=last_hidden_states.device), sequence_lengths]
def parse_lyrics(lyrics_lines: list[str]) -> list[str]:
"""Parse lyrics file and return content, ignoring timestamps."""
parsed_lyrics = []
# Pattern to identify and remove timestamps like [00:12.34]
timestamp_pattern = re.compile(r'\[\d+:\d+(?:\.\d+)?\]')
for line in lyrics_lines:
# Remove all timestamps from the line
content = timestamp_pattern.sub('', line).strip()
if content: # Skip empty lyrics
parsed_lyrics.append(content)
return parsed_lyrics
def get_detailed_instruct(task_description: str, query: str) -> str:
"""Generate query format with instruction."""
return f'Instruct: {task_description}\nQuery: {query}'
def load_descriptions(scenes_dir: str) -> dict[str, str]:
"""Load all scene descriptions and remove newlines."""
descriptions = {}
if not os.path.exists(scenes_dir):
raise FileNotFoundError(f"Scenes directory not found: {scenes_dir}")
video_folders = sorted([f for f in os.listdir(scenes_dir)
if os.path.isdir(os.path.join(scenes_dir, f))])
print(f"[INFO] Found {len(video_folders)} video folder(s) to process")
for video_folder in video_folders:
descriptions_path = os.path.join(scenes_dir, video_folder, "descriptions.json")
if not os.path.exists(descriptions_path):
print(f"[WARNING] Description file {descriptions_path} not found, skipping")
continue
try:
with open(descriptions_path, "r", encoding="utf-8") as f:
folder_descriptions = json.load(f)
# Remove newlines from descriptions
for key, value in folder_descriptions.items():
descriptions[key] = value.replace('\n', '')
except json.JSONDecodeError as e:
print(f"[WARNING] Failed to parse {descriptions_path}: {e}")
continue
return descriptions
if __name__ == "__main__":
# Load tokenizer and model
print("[INFO] Loading tokenizer and model...")
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME, padding_side='left')
# We recommend enabling flash_attention_2 for better acceleration and memory saving.
model = AutoModel.from_pretrained(
MODEL_NAME,
attn_implementation="flash_attention_2",
torch_dtype=torch.float16
).cuda()
# Read and parse lyrics
with open(LRC_PATH, "r", encoding="utf-8") as f:
lyrics = f.readlines()
parsed_lyrics = parse_lyrics(lyrics)
# Load scene descriptions
descriptions = load_descriptions(SCENES_DIR)
if not descriptions:
raise ValueError("No scene descriptions loaded")
# Prepare queries and documents
queries = [get_detailed_instruct(TASK_INSTRUCTION, lyric) for lyric in parsed_lyrics]
documents = list(descriptions.values())
scene_keys = list(descriptions.keys())
print(f"[INFO] Processing {len(parsed_lyrics)} lyrics and {len(documents)} scene descriptions")
# Encode queries in batches
print("[INFO] Encoding queries...")
query_embeddings_list = []
for i in range(0, len(queries), BATCH_SIZE):
batch_queries = queries[i:i + BATCH_SIZE]
query_batch = tokenizer(
batch_queries,
padding=True,
truncation=True,
max_length=MAX_LENGTH,
return_tensors="pt",
)
query_batch.to(model.device)
with torch.no_grad():
query_outputs = model(**query_batch)
batch_embeddings = last_token_pool(query_outputs.last_hidden_state, query_batch['attention_mask'])
batch_embeddings = F.normalize(batch_embeddings, p=2, dim=1)
query_embeddings_list.append(batch_embeddings.cpu())
# Clear cache to free memory
del query_batch, query_outputs
torch.cuda.empty_cache()
query_embeddings = torch.cat(query_embeddings_list, dim=0).cuda()
# Try to load embeddings from cache
loaded_cache = False
if os.path.exists(SCENE_EMBEDDINGS_PATH):
print(f"[INFO] Loading document embeddings from {SCENE_EMBEDDINGS_PATH}...")
try:
cache_data = torch.load(SCENE_EMBEDDINGS_PATH)
if cache_data['keys'] == scene_keys:
document_embeddings = cache_data['embeddings'].to(model.device)
print("[INFO] Embeddings loaded successfully.")
loaded_cache = True
else:
print("[WARNING] Cached keys do not match current scene keys. Recomputing embeddings.")
except Exception as e:
print(f"[WARNING] Failed to load embeddings cache: {e}. Recomputing embeddings.")
if not loaded_cache:
# Encode documents in batches
print(f"[INFO] Encoding documents in batches of {BATCH_SIZE}...")
document_embeddings_list = []
total_batches = (len(documents) + BATCH_SIZE - 1) // BATCH_SIZE
for i in range(0, len(documents), BATCH_SIZE):
batch_docs = documents[i:i + BATCH_SIZE]
batch_num = i // BATCH_SIZE + 1
if batch_num % 10 == 0 or batch_num == total_batches:
print(f"[INFO] Processing document batch {batch_num}/{total_batches}")
doc_batch = tokenizer(
batch_docs,
padding=True,
truncation=True,
max_length=MAX_LENGTH,
return_tensors="pt",
)
doc_batch.to(model.device)
with torch.no_grad():
doc_outputs = model(**doc_batch)
batch_embeddings = last_token_pool(doc_outputs.last_hidden_state, doc_batch['attention_mask'])
batch_embeddings = F.normalize(batch_embeddings, p=2, dim=1)
document_embeddings_list.append(batch_embeddings.cpu())
# Clear cache to free memory
del doc_batch, doc_outputs
torch.cuda.empty_cache()
document_embeddings = torch.cat(document_embeddings_list, dim=0).cuda()
# Save embeddings
print(f"[INFO] Saving document embeddings to {SCENE_EMBEDDINGS_PATH}...")
torch.save({'keys': scene_keys, 'embeddings': document_embeddings.cpu()}, SCENE_EMBEDDINGS_PATH)
# Calculate similarity scores
scores = query_embeddings @ document_embeddings.T
# Get best matches
top_scores, top_indices = torch.topk(scores, k=min(TOP_K, len(documents)), dim=1)
# Save results
with open(BEST_MATCHES_PATH, "w", encoding="utf-8") as f:
for i, lyric in enumerate(parsed_lyrics):
f.write(f"Lyric: {lyric}\n")
for rank in range(top_indices.shape[1]):
idx = top_indices[i, rank].item()
score = top_scores[i, rank].item()
f.write(f" Match {rank + 1}: {scene_keys[idx]} (Score: {score:.4f})\n")
f.write("\n")
print(f"[DONE] Top {TOP_K} matches for each lyric saved to {BEST_MATCHES_PATH}")