Yuvvan Talreja commited on
Commit
dd0d525
·
1 Parent(s): b398b47
Files changed (2) hide show
  1. Dockerfile +20 -9
  2. fastapi_app.py +58 -57
Dockerfile CHANGED
@@ -10,18 +10,22 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
10
  && apt-get clean \
11
  && rm -rf /var/lib/apt/lists/*
12
 
 
 
 
13
  # Set environment variables for controlling cache locations
14
- ENV TRANSFORMERS_CACHE=/app/.cache/huggingface/transformers
15
- ENV HF_HOME=/app/.cache/huggingface
16
- ENV HF_DATASETS_CACHE=/app/.cache/huggingface/datasets
17
- ENV SENTENCE_TRANSFORMERS_HOME=/app/.cache/torch/sentence_transformers
18
- ENV TORCH_HOME=/app/.cache/torch
19
 
20
  # Create cache directories with proper permissions
21
- RUN mkdir -p /app/.cache/huggingface/transformers
22
- RUN mkdir -p /app/.cache/huggingface/datasets
23
- RUN mkdir -p /app/.cache/torch/sentence_transformers
24
- RUN mkdir -p /app/.cache/torch/hub
 
25
 
26
  # Copy requirements first for better caching
27
  COPY requirements.txt .
@@ -42,6 +46,13 @@ RUN python -m nltk.downloader punkt wordnet
42
  # Copy the rest of the application
43
  COPY . .
44
 
 
 
 
 
 
 
 
45
  # Expose the port Hugging Face Spaces expects (7860)
46
  EXPOSE 7860
47
 
 
10
  && apt-get clean \
11
  && rm -rf /var/lib/apt/lists/*
12
 
13
+ # Create a non-root user to run the application
14
+ RUN useradd -m appuser
15
+
16
  # Set environment variables for controlling cache locations
17
+ ENV TRANSFORMERS_CACHE=/tmp/.cache/huggingface/transformers
18
+ ENV HF_HOME=/tmp/.cache/huggingface
19
+ ENV HF_DATASETS_CACHE=/tmp/.cache/huggingface/datasets
20
+ ENV SENTENCE_TRANSFORMERS_HOME=/tmp/.cache/torch/sentence_transformers
21
+ ENV TORCH_HOME=/tmp/.cache/torch
22
 
23
  # Create cache directories with proper permissions
24
+ RUN mkdir -p /tmp/.cache/huggingface/transformers && \
25
+ mkdir -p /tmp/.cache/huggingface/datasets && \
26
+ mkdir -p /tmp/.cache/torch/sentence_transformers && \
27
+ mkdir -p /tmp/.cache/torch/hub && \
28
+ chmod -R 777 /tmp/.cache
29
 
30
  # Copy requirements first for better caching
31
  COPY requirements.txt .
 
46
  # Copy the rest of the application
47
  COPY . .
48
 
49
+ # Make sure the non-root user can access the application files
50
+ RUN chown -R appuser:appuser /app
51
+ RUN chmod -R 755 /app
52
+
53
+ # Switch to non-root user
54
+ USER appuser
55
+
56
  # Expose the port Hugging Face Spaces expects (7860)
57
  EXPOSE 7860
58
 
fastapi_app.py CHANGED
@@ -8,7 +8,7 @@ import base64
8
  import numpy as np
9
  import torch
10
  import time
11
- from transformers import pipeline, AutoProcessor, AutoModelForSpeechSeq2Seq
12
  import threading
13
  import json
14
  import logging
@@ -16,32 +16,11 @@ from pydantic import BaseModel
16
  from typing import Optional, Dict, Any, List
17
  from topic_segmenter import TopicSegmenter
18
  from middleware import HTTPSProxyMiddleware
19
- import tempfile
20
 
21
  # Configure logging
22
  logging.basicConfig(level=logging.INFO)
23
  logger = logging.getLogger(__name__)
24
 
25
- # Set environment variables for cache directories to writable locations
26
- os.environ["TRANSFORMERS_CACHE"] = "/tmp/transformers_cache"
27
- os.environ["HF_HOME"] = "/tmp/hf_home"
28
- os.environ["XDG_CACHE_HOME"] = "/tmp/xdg_cache"
29
- os.environ["NLTK_DATA"] = "/tmp/nltk_data"
30
-
31
- # Create directories
32
- for directory in ["/tmp/transformers_cache", "/tmp/hf_home", "/tmp/xdg_cache", "/tmp/nltk_data"]:
33
- os.makedirs(directory, exist_ok=True)
34
-
35
- # Try to download NLTK data to the tmp directory
36
- try:
37
- import nltk
38
- nltk.data.path.append("/tmp/nltk_data")
39
- nltk.download('punkt', download_dir="/tmp/nltk_data", quiet=True)
40
- nltk.download('wordnet', download_dir="/tmp/nltk_data", quiet=True)
41
- nltk.download('omw-1.4', download_dir="/tmp/nltk_data", quiet=True)
42
- except Exception as e:
43
- logger.warning(f"NLTK download issue: {str(e)}")
44
-
45
  app = FastAPI(title="Flowify", description="Audio transcription and topic segmentation API")
46
 
47
  # Add custom middleware for HTTPS handling
@@ -95,49 +74,68 @@ def get_model(model_name):
95
  if model_name not in models:
96
  logging.info(f"Loading model: {model_name}")
97
 
98
- try:
99
- # Direct approach to load the model with correct parameters
100
- processor = AutoProcessor.from_pretrained(
101
- model_name,
102
- cache_dir="/tmp/transformers_cache",
103
- local_files_only=False
104
- )
 
 
105
 
106
- model = AutoModelForSpeechSeq2Seq.from_pretrained(
107
- model_name,
108
- torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,
109
- cache_dir="/tmp/transformers_cache",
110
- local_files_only=False
111
- )
112
 
 
113
  models[model_name] = pipeline(
114
- "automatic-speech-recognition",
115
- model=model,
116
- tokenizer=processor.tokenizer,
117
- feature_extractor=processor.feature_extractor,
118
  chunk_length_s=30,
119
- stride_length_s=5
 
 
 
 
 
120
  )
121
- logging.info(f"Model {model_name} loaded successfully")
122
-
123
  except Exception as e:
124
- logging.error(f"Error loading model: {str(e)}")
125
-
126
- # Fallback - try using just the pipeline which sometimes works better
127
  try:
128
- logging.info(f"Trying fallback method for {model_name}")
 
 
 
 
 
 
 
 
 
 
 
129
  models[model_name] = pipeline(
130
- "automatic-speech-recognition",
131
- model=model_name,
132
- torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,
 
133
  chunk_length_s=30,
134
- stride_length_s=5
 
 
 
 
 
135
  )
136
- logging.info(f"Model {model_name} loaded successfully using fallback")
137
  except Exception as e2:
138
- logging.error(f"Fallback loading also failed: {str(e2)}")
139
- raise RuntimeError(f"Failed to load model {model_name}: {str(e)} | Fallback error: {str(e2)}")
140
-
141
  return models[model_name]
142
 
143
  def process_audio(audio_data, sample_rate=16000, model_name="openai/whisper-base"):
@@ -158,10 +156,14 @@ def process_audio(audio_data, sample_rate=16000, model_name="openai/whisper-base
158
 
159
  logging.info(f"Processing chunk {i+1}/{total_chunks}")
160
 
 
161
  result = model(
162
  chunk,
163
  return_timestamps=True,
164
- generate_kwargs={"language": "english", "task": "transcribe"}
 
 
 
165
  )
166
 
167
  if "chunks" in result and len(result["chunks"]) > 0:
@@ -212,8 +214,7 @@ async def analyze_topics(data: TranscriptData):
212
  min_segment_size=2,
213
  topic_similarity_threshold=0.25,
214
  max_topics=6,
215
- hierarchical_threshold=0.6,
216
- model_name="all-MiniLM-L6-v2"
217
  )
218
 
219
  segments, topic_mappings, topic_history, topic_hierarchies = segmenter.segment_transcript(transcript)
 
8
  import numpy as np
9
  import torch
10
  import time
11
+ from transformers import pipeline
12
  import threading
13
  import json
14
  import logging
 
16
  from typing import Optional, Dict, Any, List
17
  from topic_segmenter import TopicSegmenter
18
  from middleware import HTTPSProxyMiddleware
 
19
 
20
  # Configure logging
21
  logging.basicConfig(level=logging.INFO)
22
  logger = logging.getLogger(__name__)
23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
24
  app = FastAPI(title="Flowify", description="Audio transcription and topic segmentation API")
25
 
26
  # Add custom middleware for HTTPS handling
 
74
  if model_name not in models:
75
  logging.info(f"Loading model: {model_name}")
76
 
77
+ hf_model_name = model_name
78
+ if model_name.startswith('Xenova/'):
79
+ base_name = model_name.split('/')[-1]
80
+ if '.en' in base_name:
81
+ size = base_name.split('.')[0].replace('whisper-', '')
82
+ hf_model_name = f"openai/whisper-{size}"
83
+ else:
84
+ size = base_name.replace('whisper-', '')
85
+ hf_model_name = f"openai/whisper-{size}"
86
 
87
+ logging.info(f"Converting Xenova model {model_name} to {hf_model_name}")
88
+
89
+ try:
90
+ # Load the model using TensorFlow compatibility mode
91
+ import tensorflow as tf
92
+ logging.info(f"TensorFlow version: {tf.__version__}")
93
 
94
+ # Configure pipeline with correct parameters to avoid conflicts
95
  models[model_name] = pipeline(
96
+ "automatic-speech-recognition",
97
+ model=hf_model_name,
 
 
98
  chunk_length_s=30,
99
+ stride_length_s=5,
100
+ framework="tf", # Explicitly use TensorFlow
101
+ model_kwargs={
102
+ "attention_mask": True, # Explicitly set attention mask
103
+ "use_cache": True,
104
+ }
105
  )
106
+ logging.info(f"Model {model_name} loaded successfully with TensorFlow")
 
107
  except Exception as e:
108
+ logging.error(f"Error loading model with TensorFlow: {str(e)}")
 
 
109
  try:
110
+ # Fallback to PyTorch with more specific configurations
111
+ from transformers import AutoProcessor, AutoModelForSpeechSeq2Seq
112
+
113
+ # First load processor and model separately to customize configs
114
+ processor = AutoProcessor.from_pretrained(hf_model_name)
115
+ model = AutoModelForSpeechSeq2Seq.from_pretrained(
116
+ hf_model_name,
117
+ use_cache=True,
118
+ attention_mask=True
119
+ )
120
+
121
+ # Create the pipeline with the initialized model and processor
122
  models[model_name] = pipeline(
123
+ "automatic-speech-recognition",
124
+ model=model,
125
+ tokenizer=processor,
126
+ feature_extractor=processor,
127
  chunk_length_s=30,
128
+ stride_length_s=5,
129
+ generate_kwargs={
130
+ "task": "transcribe",
131
+ "language": "english",
132
+ # Don't set forced_decoder_ids here to avoid conflict
133
+ }
134
  )
135
+ logging.info(f"Model {model_name} loaded successfully with PyTorch")
136
  except Exception as e2:
137
+ logging.error(f"Failed to load model {model_name}: {str(e2)}")
138
+ raise RuntimeError(f"Failed to load model {model_name}: {str(e2)}")
 
139
  return models[model_name]
140
 
141
  def process_audio(audio_data, sample_rate=16000, model_name="openai/whisper-base"):
 
156
 
157
  logging.info(f"Processing chunk {i+1}/{total_chunks}")
158
 
159
+ # Use consistent parameters for model inference
160
  result = model(
161
  chunk,
162
  return_timestamps=True,
163
+ generate_kwargs={
164
+ "task": "transcribe",
165
+ "language": "english"
166
+ }
167
  )
168
 
169
  if "chunks" in result and len(result["chunks"]) > 0:
 
214
  min_segment_size=2,
215
  topic_similarity_threshold=0.25,
216
  max_topics=6,
217
+ hierarchical_threshold=0.6
 
218
  )
219
 
220
  segments, topic_mappings, topic_history, topic_hierarchies = segmenter.segment_transcript(transcript)