File size: 2,376 Bytes
de1e3fc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Loading with skip_embeddings stores images only; embed-all generates embeddings later."""
from pathlib import Path

import pytest
from sqlalchemy import func, select

from app.db import SessionLocal
from app.models import BreedPrediction, Dataset, Embedding, Picture
from app.services.batch_loader import load_dataset
from app.services.datasets import embed_all
from scripts.make_sample_images import make_image
from scripts.prepare_test_data import prepare


@pytest.fixture
def db():
    session = SessionLocal()
    try:
        yield session
    finally:
        session.rollback()
        session.close()


def _build_input(tmp: Path) -> Path:
    root = tmp / "input"
    for folder in ("dogA", "dogB"):
        d = root / folder
        d.mkdir(parents=True)
        for i in range(2):
            (d / f"img{i}.jpg").write_bytes(make_image(abs(hash((folder, i))) % 1000))
    return root


def test_load_skips_embeddings_then_embed_all_generates(tmp_path, db):
    root = _build_input(tmp_path)
    prepare(root, root, holdout=1, seed=1)

    result = load_dataset(
        db, folder=root, dataset_type="known", name="StorageOnly", description=None,
        csv_path=root / "known_dogs.csv", skip_embeddings=True,
    )
    dataset = db.get(Dataset, result.dataset_id)

    # Pictures were stored, but no embeddings or breed predictions were generated.
    pic_ids = db.execute(
        select(Picture.id).where(Picture.subject_type == "known")
    ).scalars().all()
    assert len(pic_ids) == result.images_processed >= 1
    assert db.execute(select(func.count()).select_from(Embedding)).scalar_one() == 0
    assert db.execute(select(func.count()).select_from(BreedPrediction)).scalar_one() == 0

    # embed-all backfills embeddings for the stored images.
    summary = embed_all(db, dataset)
    assert summary["embedded"] == summary["total_pictures"] == len(pic_ids)
    assert db.execute(select(func.count()).select_from(Embedding)).scalar_one() == len(pic_ids)


def test_default_load_still_embeds(tmp_path, db):
    root = _build_input(tmp_path)
    prepare(root, root, holdout=1, seed=1)
    load_dataset(
        db, folder=root, dataset_type="known", name="WithEmb", description=None,
        csv_path=root / "known_dogs.csv",  # skip_embeddings defaults False
    )
    assert db.execute(select(func.count()).select_from(Embedding)).scalar_one() >= 1