Harsh-P/knowledge_graph
0
1import whisper2from pyannote.audio import Pipeline3from datetime import timedelta4import torch5 6import nest_asyncio7import os8from llama_index.llms.groq import Groq9from llama_index.embeddings.huggingface import HuggingFaceEmbedding10from llama_index.core import Settings11from llama_index.core.llms import ChatMessage12 13def get_transcript(audio):14 # === CONFIGURATION ===15 #AUDIO_FILE = "/content/sample audio.mp3"16 HF_TOKEN = os.environ['TRANSCRIPT_TOKEN']17 18 # === STEP 1: Transcribe with Whisper ===19 print("๐ Transcribing with Whisper...")20 whisper_model = whisper.load_model("large")21 whisper_result = whisper_model.transcribe(audio)22 23 # === STEP 2: Perform Speaker Diarization ===24 print("๐ Performing Speaker Diarization...")25 pipeline = Pipeline.from_pretrained(26 "pyannote/speaker-diarization-3.1",27 use_auth_token=HF_TOKEN28 )29 # pipeline.to(torch.device("cuda"))30 31 diarization = pipeline(AUDIO_FILE)32 33 # === STEP 3: Match Whisper segments to diarization segments ===34 # === STEP 3: Assign speaker to each Whisper segment based on diarization ===35 print("๐ Assigning speaker labels to Whisper segments (using midpoint matching)...")36 37 speaker_segments = []38 for seg in whisper_result["segments"]:39 seg_start = seg["start"]40 seg_end = seg["end"]41 seg_mid = (seg_start + seg_end) / 242 assigned_speaker = "Unknown"43 44 for turn, _, speaker in diarization.itertracks(yield_label=True):45 if turn.start <= seg_mid <= turn.end:46 assigned_speaker = speaker47 break48 49 speaker_segments.append({50 "speaker": assigned_speaker,51 "start": seg_start,52 "end": seg_end,53 "text": seg["text"].strip()54 })55 56 57 58 # === STEP 4: Format output ===59 # === STEP 4: Format output ===60 print("\nโ
Final Speaker-Attributed Transcript:\n")61 62 def format_time(seconds):63 return str(timedelta(seconds=float(f"{seconds:.3f}")))64 65 66 transcript=""67 for seg in speaker_segments:68 if seg["text"]:69 start = format_time(seg["start"])70 end = format_time(seg["end"])71 output_line = f"[{seg['speaker']}] ({start} - {end}): {seg['text']}"72 transcript+=output_line+"\n"73 #print(output_line)74 return transcript75 76 77 78 # # === OPTIONAL: Save to file ===79 # with open("speaker_transcript.txt", "w") as f:80 # for seg in aligned_segments:81 # if seg["text"]:82 # start = format_time(seg["start"])83 # end = format_time(seg["end"])84 # f.write(f"[{seg['speaker']}] ({start} - {end}): {seg['text']}\n")85 86 87def get_triplets(transcript):88 # โ
Apply Nest AsyncIO89 nest_asyncio.apply()90 91 # โ
Set Groq API key92 os.environ["GROQ_API_KEY"] = "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"93 94 # โ
Initialize LLM and Embeddings95 llm = Groq(model="llama3-8b-8192")96 embed_model = HuggingFaceEmbedding(model_name="BAAI/bge-small-en-v1.5")97 98 # โ
Configure global settings99 Settings.llm = llm100 Settings.embed_model = embed_model101 102 # โ
Ask Llama3 to extract triplets103 prompt = f"""104 Extract the structured triplets from the following meeting transcript. Format your output as a list of (subject, predicate, object) triplets like below.105 [106 ("SPEAKER_01", "asks", "Sid about bug fixes from yesterday"),107 ("SPEAKER_00", "says", "Crash on profile load fixed, testing edge cases"),108 ("SPEAKER_01", "asks", "if analytics integration was reviewed"),109 ("SPEAKER_00", "plans", "to plug in events after lunch"),110 ("SPEAKER_01", "wants", "build ready by 4 to test on phone"),111 ("SPEAKER_00", "promises", "build will be ready before 3.30"),112 ("SPEAKER_01", "reminds", "launch is in 3 days"),113 ("SPEAKER_00", "confirms", "almost done, just final checks")114 ]115 116 117 Transcript:118 {transcript}119 """120 121 response = llm.complete(prompt)122 123 # Extract the part between the square brackets124 match = re.search(r'\[\s*(\([^\]]+\))\s*\]', response, re.DOTALL)125 triplet=[]126 if match:127 triplet_str = "[" + match.group(1) + "]"128 triplet_list = ast.literal_eval(triplet_str)129 #print(triplet_list)130 else:131 print("No triplet list found.")132 triplets=triplet_list133 134 return triplets135 136def push_graph(triplets):137 from neo4j import GraphDatabase138 139 # === Your Neo4j connection info ===140 URI = "neo4j+s://a12094a4.databases.neo4j.io" # or from Aura dashboard141 USERNAME = "neo4j"142 PASSWORD = "qI4aBidXAhd7RYYBXQ-apAePWUOtytJ_dsfZA6LK__k"143 144 # # === Triplets from earlier ===145 # triplets = [146 # ("SPEAKER_01", "asks", "Sid about bug fixes from yesterday"),147 # ("SPEAKER_00", "says", "Crash on profile load fixed, testing edge cases"),148 # ("SPEAKER_01", "asks", "if analytics integration was reviewed"),149 # ("SPEAKER_00", "plans", "to plug in events after lunch"),150 # ("SPEAKER_01", "wants", "build ready by 4 to test on phone"),151 # ("SPEAKER_00", "promises", "build will be ready before 3.30"),152 # ("SPEAKER_01", "reminds", "launch is in 3 days"),153 # ("SPEAKER_00", "confirms", "almost done, just final checks")154 # ]155 156 # === Neo4j interaction ===157 class KnowledgeGraphUploader:158 def __init__(self, uri, user, password):159 self.driver = GraphDatabase.driver(uri, auth=(user, password))160 161 def close(self):162 self.driver.close()163 164 def upload_triplets(self, triplets):165 with self.driver.session() as session:166 for subj, pred, obj in triplets:167 session.write_transaction(self._create_relationship, subj, pred, obj)168 169 @staticmethod170 def _create_relationship(tx, subj, pred, obj):171 query = (172 "MERGE (a:Entity {name: $subj}) "173 "MERGE (b:Entity {name: $obj}) "174 "MERGE (a)-[r:RELATION {type: $pred}]->(b)"175 )176 tx.run(query, subj=subj, pred=pred, obj=obj)177 178 179 # === Run it ===180 kg_uploader = KnowledgeGraphUploader(URI, USERNAME, PASSWORD)181 kg_uploader.upload_triplets(triplets)182 kg_uploader.close()183 184 #print("โ
Knowledge graph uploaded to Neo4j!")185 186def generate_knowledge_graph(audio):187 transcript=get_transcript(audio)188 triplets=get_triplets(transcript)189 push_graph(triplets)190 191from fastapi import FastAPI, UploadFile, File, HTTPException192from fastapi.responses import JSONResponse193import shutil194import os195import uuid196 197app = FastAPI()198 199# Your knowledge graph function (plug your real function here)200# from your_module import generate_knowledge_graph201 202@app.post("/generate-graph/")203async def generate_graph_from_audio(file: UploadFile = File(...)):204 # Validate MP3 format205 if not file.filename.lower().endswith('.mp3'):206 raise HTTPException(status_code=400, detail="Only MP3 files are supported.")207 208 # Save the MP3 file temporarily209 temp_dir = "temp_uploads"210 os.makedirs(temp_dir, exist_ok=True)211 temp_filename = f"{uuid.uuid4()}.mp3"212 temp_path = os.path.join(temp_dir, temp_filename)213 214 with open(temp_path, "wb") as buffer:215 shutil.copyfileobj(file.file, buffer)216 217 try:218 # Trigger knowledge graph creation (stored in Neo4j)219 generate_knowledge_graph(temp_path)220 221 # Return a basic confirmation response222 return JSONResponse(content={223 "status": "success",224 "message": "Knowledge graph generated and uploaded to Neo4j successfully."225 })226 227 except Exception as e:228 raise HTTPException(status_code=500, detail=str(e))229 230 finally:231 # Clean up the temp file232 if os.path.exists(temp_path):233 os.remove(temp_path)234 