0. 图片与文本提示¶
代码会加载预训练 CLIP,但页面重点不是模型下载,而是读懂图文匹配矩阵:行是图片,列是文本,数值越大表示越匹配。
In [3]:
# 真实图片与文本提示:计算每张图片更匹配哪一句描述。
model_id = "openai/clip-vit-base-patch32"
with contextlib.redirect_stdout(io.StringIO()), contextlib.redirect_stderr(io.StringIO()):
processor = CLIPProcessor.from_pretrained(model_id, use_fast=True)
clip_model = CLIPModel.from_pretrained(model_id)
clip_model.eval()
clip_images = [
Image.fromarray(load_sample_image("china.jpg")),
Image.fromarray(load_sample_image("flower.jpg")),
]
image_names = ["china.jpg", "flower.jpg"]
image_labels = ["湖边寺庙图", "红色花朵图"]
text_prompts = [
"a photo of a Chinese temple by a lake",
"a close-up photo of a red flower",
"a photo of a taxi cab",
"a photo of a city street at night",
]
prompt_labels = ["湖边寺庙", "红色花朵", "出租车", "夜晚街道"]
inputs = processor(text=text_prompts, images=clip_images, return_tensors="pt", padding=True)
with torch.no_grad():
outputs = clip_model(**inputs)
logits = outputs.logits_per_image
probs = logits.softmax(dim=1)
targets = torch.tensor([0, 1])
clip_loss = F.cross_entropy(logits[:, :2], targets).item()
clip_prob_df = pd.DataFrame(probs.numpy(), index=image_labels, columns=prompt_labels)
display(clip_prob_df.round(3))
print("图文对比损失:", round(float(clip_loss), 4))
| 湖边寺庙 | 红色花朵 | 出租车 | 夜晚街道 | |
|---|---|---|---|---|
| 湖边寺庙图 | 1.0 | 0.000 | 0.000 | 0.0 |
| 红色花朵图 | 0.0 | 0.994 | 0.006 | 0.0 |
图文对比损失: 0.0001
In [4]:
# 绘制真实图片和图文匹配概率矩阵。
sim_matrix = probs.numpy()
fig = plt.figure(figsize=(11.0, 5.4))
gs = fig.add_gridspec(1, 3, width_ratios=[1.0, 1.0, 2.2], wspace=0.35)
for idx, image in enumerate(clip_images):
ax_img = fig.add_subplot(gs[0, idx])
ax_img.imshow(image)
ax_img.set_title(image_labels[idx], fontweight="bold")
ax_img.set_xticks([])
ax_img.set_yticks([])
ax = fig.add_subplot(gs[0, 2])
im = ax.imshow(sim_matrix, cmap="YlGnBu", vmin=0, vmax=1)
ax.set_xticks(range(len(text_prompts)), prompt_labels, rotation=30, ha="right")
ax.set_yticks(range(len(image_names)), image_labels)
for i in range(len(image_names)):
best = int(np.argmax(sim_matrix[i]))
for j in range(len(text_prompts)):
value = sim_matrix[i, j]
ax.text(j, i, f"{value:.2f}", ha="center", va="center", color="#0f172a", fontweight="bold" if j == best else "normal")
if j == best:
ax.add_patch(plt.Rectangle((j - 0.5, i - 0.5), 1, 1, fill=False, edgecolor="#0f172a", linewidth=2.2))
ax.set_title("图文匹配概率", loc="left", fontsize=14, fontweight="bold", color="#0f172a")
fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
plt.tight_layout()
plt.show()