-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
127 lines (109 loc) · 4.91 KB
/
Copy pathmain.py
File metadata and controls
127 lines (109 loc) · 4.91 KB
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
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
import argparse
import os
import sys
import json
import logging
from typing import Dict
from functools import lru_cache
import torch
import google.generativeai as genai
from fastapi import FastAPI
import uvicorn
# Import project modules (assuming paths are set)
from models.transformer import CosmicTransformer
from models.diffusion import AudioDiffusion # For hybrid use if needed
from models.cache import get_cached_model
from scripts.train import train_model
from api.app import create_app # FastAPI app factory
from data.loader import AudioTextDataset # For training
# Setup private logging (in-memory, no external exposure)
logging.basicConfig(level=os.environ.get('LOG_LEVEL', 'INFO'), stream=sys.stdout)
logger = logging.getLogger(__name__)
# Configure Gemini API once at module level
try:
if 'GEMINI_API_KEY' in os.environ:
genai.configure(api_key=os.environ['GEMINI_API_KEY'])
except Exception as e:
logger.warning(f"Failed to configure Gemini API: {e}")
# Gemini bootstrap function with caching
@lru_cache(maxsize=128)
def _cached_gemini_call(gemini_prompt: str) -> str:
"""Cached Gemini API call to avoid redundant requests."""
try:
model = genai.GenerativeModel('gemini-pro')
response = model.generate_content(gemini_prompt)
return response.text
except Exception as e:
logger.warning(f"Gemini API call failed: {e}")
raise
def refine_prompt_with_gemini(raw_prompt: str, neural_input: Dict) -> str:
"""Use Gemini-Pro to refine prompt with neural data integration."""
try:
# Calculated neural modulation equation
if 'heart_rate' in neural_input:
hr = neural_input['heart_rate']
tempo_adjust = 60 + (hr / 150) * 120 # Scale to 60-180 BPM
neural_desc = f" with tempo around {int(tempo_adjust)} BPM reflecting elevated energy"
else:
neural_desc = ""
gemini_prompt = f"Enhance this music prompt for AI generation: '{raw_prompt}'. Incorporate emotional cues{neural_desc}. Output detailed description."
return _cached_gemini_call(gemini_prompt)
except KeyError:
logger.error("GEMINI_API_KEY not set.")
sys.exit(1)
except Exception as e:
logger.warning(f"Gemini refinement failed: {e}. Using raw prompt.")
return raw_prompt
# Inference function
def generate_music(prompt: str, neural_input: Dict, output_file: str) -> str:
"""Generate audio using refined prompt and model."""
refined_prompt = refine_prompt_with_gemini(prompt, neural_input)
logger.info(f"Refined prompt: {refined_prompt}")
model = get_cached_model(os.environ.get('MODEL_CHECKPOINT', 'models/checkpoint.pth'))
audio = model.generate(refined_prompt, neural_input) # Assume generate returns torch tensor
# Save audio (using torchaudio or ffmpeg)
import torchaudio
torchaudio.save(output_file, audio, sample_rate=44100)
return output_file
# Main CLI parser
def main():
parser = argparse.ArgumentParser(description="Cosmic Composer CLI: Bootstrap AI music framework with Gemini-Pro integration.")
subparsers = parser.add_subparsers(dest='mode', required=True)
# Train subcommand
train_parser = subparsers.add_parser('train', help='Train the model')
train_parser.add_argument('--dataset', required=True, help='Path to dataset')
train_parser.add_argument('--epochs', type=int, default=50)
train_parser.add_argument('--batch-size', type=int, default=16)
# Generate subcommand
gen_parser = subparsers.add_parser('generate', help='Generate music from prompt')
gen_parser.add_argument('--prompt', required=True, help='Text prompt for music')
gen_parser.add_argument('--neural-input', type=str, default='{}', help='JSON neural input, e.g. {"heart_rate": 90}')
gen_parser.add_argument('--output-file', default='output.wav', help='Output audio file')
# API subcommand
api_parser = subparsers.add_parser('api', help='Start FastAPI server')
api_parser.add_argument('--host', default='0.0.0.0')
api_parser.add_argument('--port', type=int, default=8000)
# Test subcommand
test_parser = subparsers.add_parser('test', help='Run tests')
args = parser.parse_args()
if args.mode == 'train':
dataset = AudioTextDataset(args.dataset)
train_model(
dataset=dataset,
epochs=args.epochs,
batch_size=args.batch_size
)
elif args.mode == 'generate':
neural_input = json.loads(args.neural_input)
output = generate_music(args.prompt, neural_input, args.output_file)
logger.info(f"Generated audio saved to {output}")
elif args.mode == 'api':
app: FastAPI = create_app() # Factory to include models/routes
uvicorn.run(app, host=args.host, port=args.port)
elif args.mode == 'test':
import pytest
sys.exit(pytest.main(['tests/']))
else:
parser.print_help()
if __name__ == '__main__':
main()