Team Ai
Apppublic

Harsh-P/knowledge_graph

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py234 linesDownload Raw Back to root
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