cck-0702's picture
Clean commit without binary (image) files
c8c00f0
Raw
History Blame Contribute Delete
6.33 kB
# ============================================
# DefectRAG — DefectDiffu Edition
# ============================================
# Retrieves in-context defect examples from ChromaDB and composes
# DefectDiffu text prompts (c_d) from retrieved knowledge.
import chromadb
from sentence_transformers import SentenceTransformer
import json
from pathlib import Path
from typing import List, Dict, Optional
class DefectRAG:
"""
Retrieval-Augmented Generation for defect examples.
Now includes helper to build DefectDiffu defect prompts (c_d) from RAG results.
"""
def __init__(
self,
db_path: str = "data/defect_db",
collection: str = "defect_patches",
model: str = "all-MiniLM-L6-v2"
):
self.db_path = db_path
self.collection_name = collection
self._client = None
self._collection = None
self._encoder = None
@property
def client(self):
if self._client is None:
self._client = chromadb.PersistentClient(path=self.db_path)
return self._client
@property
def collection(self):
if self._collection is None:
self._collection = self.client.get_collection(self.collection_name)
return self._collection
@property
def encoder(self):
if self._encoder is None:
self._encoder = SentenceTransformer('all-MiniLM-L6-v2')
return self._encoder
def _build_where_filter(
self,
commercial_only: bool = True,
domain_filter: Optional[str] = None
) -> Optional[Dict]:
conditions = []
if commercial_only:
conditions.append({"commercial_ok": True})
if domain_filter and domain_filter != "general":
conditions.append({"domain": domain_filter})
if len(conditions) == 0:
return None
elif len(conditions) == 1:
return conditions[0]
else:
return {"$and": conditions}
def retrieve(
self,
defect_plan,
k: int = 3,
domain_filter: Optional[str] = None,
commercial_only: bool = True
) -> List[Dict]:
"""Retrieve top-k matching defect examples."""
query_parts = [
getattr(defect_plan, 'artifact_type', defect_plan.defect_type),
defect_plan.description,
"on",
defect_plan.target_entity
]
query = " ".join(query_parts)
query_emb = self.encoder.encode(query)
where_filter = self._build_where_filter(commercial_only, domain_filter)
results = self.collection.query(
query_embeddings=[query_emb.tolist()],
n_results=k,
where=where_filter
)
examples = []
for i, meta in enumerate(results['metadatas'][0]):
paths = json.loads(meta['paths']) if isinstance(meta.get('paths'), str) else meta.get('paths', {})
examples.append({
'paths': paths,
'caption': meta.get('caption', ''),
'domain': meta.get('domain', 'unknown'),
'license': meta.get('license', 'unknown'),
'source': meta.get('source', 'unknown'),
'defect_name': meta.get('defect_name', 'unknown'),
'score': results['distances'][0][i] if results.get('distances') else None,
'metadata': {k: v for k, v in meta.items() if k not in {'paths', 'caption', 'domain', 'license', 'source', 'defect_name'}}
})
return examples
def retrieve_by_text(
self,
text: str,
k: int = 3,
domain_filter: Optional[str] = None
) -> List[Dict]:
"""Direct text search (for debugging/testing)."""
query_emb = self.encoder.encode(text)
where_filter = self._build_where_filter(True, domain_filter)
results = self.collection.query(
query_embeddings=[query_emb.tolist()],
n_results=k,
where=where_filter
)
examples = []
for i, meta in enumerate(results['metadatas'][0]):
paths = json.loads(meta['paths']) if isinstance(meta.get('paths'), str) else meta.get('paths', {})
examples.append({
'paths': paths,
'caption': meta.get('caption', ''),
'domain': meta.get('domain', 'unknown'),
'score': results['distances'][0][i] if results.get('distances') else None
})
return examples
def compose_defect_prompt(
self,
base_description: str,
examples: List[Dict],
max_examples: int = 1
) -> str:
"""
Build a DefectDiffu defect prompt (c_d) by enriching the base description
with captions from retrieved RAG examples.
Example output:
"A photo of a small transparent bubble trapped under glass, similar to
a spherical air pocket with dark meniscus ring"
"""
if not examples:
return f"A photo of {base_description}"
captions = [ex.get('caption', '') for ex in examples[:max_examples] if ex.get('caption')]
if captions:
enriched = f"{base_description}, similar to {captions[0]}"
return f"A photo of {enriched}"
return f"A photo of {base_description}"
def get_stats(self) -> Dict:
"""Get DB statistics."""
count = self.collection.count()
results = self.collection.get()
domains = {}
licenses = {}
sources = {}
for meta in results["metadatas"]:
domains[meta.get("domain", "unknown")] = domains.get(meta.get("domain"), 0) + 1
licenses[meta.get("license", "unknown")] = licenses.get(meta.get("license"), 0) + 1
sources[meta.get("source", "unknown")] = sources.get(meta.get("source"), 0) + 1
return {
"total_entries": count,
"domains": domains,
"licenses": licenses,
"sources": sources
}
# Singleton instance
_rag_instance = None
def get_rag(db_path="data/defect_db") -> DefectRAG:
"""Get or create singleton RAG instance."""
global _rag_instance
if _rag_instance is None:
_rag_instance = DefectRAG(db_path=db_path)
return _rag_instance