before changes
This commit is contained in:
@@ -0,0 +1,332 @@
|
||||
const fs = require('fs')
|
||||
const https = require('https')
|
||||
const exec = require('child_process').exec
|
||||
const AWS = require('aws-sdk')
|
||||
|
||||
const Bugsnag = require('@bugsnag/js')
|
||||
const mongoose = require('mongoose')
|
||||
const version = require('./package.json').version
|
||||
const sqs = require('../app/services/sqs')
|
||||
const s3 = require('../app/services/s3')
|
||||
const voiceCloningService = require('./voice_cloning')
|
||||
const userAudioProfileService = require('./user_audio_profile')
|
||||
|
||||
AWS.config.update({ region: 'us-west-2' })
|
||||
const sqsQueueUrl = process.env.SQS_URL
|
||||
const mongoUriDev = process.env.MONGODB_URI_DEV
|
||||
const mongoUriStaging = process.env.MONGODB_URI_STAGING
|
||||
const mongoUriProd = process.env.MONGODB_URI_PROD
|
||||
let throttleMessageFetching = true
|
||||
const APP_ENV = process.env.POTION_APP_ENV
|
||||
|
||||
const cloudFrontUrlProd = process.env.CLOUDFRONT_URL_PROD
|
||||
const cloudFrontUrlDev = process.env.CLOUDFRONT_URL_DEV
|
||||
const cloudFrontUrlStaging = process.env.CLOUDFRONT_URL_STAGING
|
||||
|
||||
const updateUrl = (str, cloudFrontUrl) => {
|
||||
const host = new URL(str).host
|
||||
return str.replace(`https://${host}`, cloudFrontUrl)
|
||||
}
|
||||
|
||||
function connectDB(dbUri, retryCount = 0) {
|
||||
return new Promise((resolve, reject) => {
|
||||
console.log('Connection Attempt : ', retryCount)
|
||||
mongoose.set('strictQuery', true)
|
||||
mongoose
|
||||
.connect(dbUri)
|
||||
.then((msg) => {
|
||||
console.log('Connected to Mongo DB !')
|
||||
resolve()
|
||||
})
|
||||
.catch((err) => {
|
||||
console.log('Failed to connect dns mongo: ', err)
|
||||
if (retryCount < 6) {
|
||||
retryCount++
|
||||
connectDB(dbUri, retryCount)
|
||||
}
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
function execShellCommand(cmd, logPath) {
|
||||
// const exec = require("child_process").exec;
|
||||
return new Promise((resolve, reject) => {
|
||||
exec(cmd, { maxBuffer: 1024 * 1000000 }, async (error, stdout, stderr) => {
|
||||
if (error) {
|
||||
console.log('Error while proccessing python command', error)
|
||||
reject(error)
|
||||
}
|
||||
// console.log('Stdout --- ', stdout)
|
||||
// console.log('Stderror --- ', stderr)
|
||||
await fs.promises.writeFile(`${logPath}/error.log`, stderr)
|
||||
await fs.promises.writeFile(`${logPath}/info.log`, stdout)
|
||||
|
||||
resolve()
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
async function getFile(waveUrl, path) {
|
||||
return new Promise((resolve) => {
|
||||
https.get(waveUrl, (res) => {
|
||||
const writeStream = fs.createWriteStream(path)
|
||||
|
||||
res.pipe(writeStream)
|
||||
|
||||
writeStream.on('finish', () => {
|
||||
writeStream.close()
|
||||
resolve()
|
||||
})
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
function pad(s) {
|
||||
while (s.length < 3) s = '0' + s // IN future we will need padding to 4
|
||||
return s
|
||||
}
|
||||
|
||||
const processQueue = () => {
|
||||
/* eslint-disable no-async-promise-executor */
|
||||
return new Promise(async (resolve, reject) => {
|
||||
try {
|
||||
const response = await sqs.fetchMessageFromSQS(sqsQueueUrl)
|
||||
|
||||
if (
|
||||
typeof response.Messages !== 'undefined' &&
|
||||
response.Messages.length > 0
|
||||
) {
|
||||
throttleMessageFetching = false
|
||||
const job = JSON.parse(response.Messages[0].Body)
|
||||
const receiptHandle = response.Messages[0].ReceiptHandle
|
||||
console.log('job===', job)
|
||||
|
||||
const { metadata, input, _id, userAudioProfileId } = job._doc
|
||||
console.log('userAudioProfileId', userAudioProfileId)
|
||||
console.log('_id', _id)
|
||||
const { env } = job
|
||||
console.log('env', env)
|
||||
|
||||
console.log('metadata------', metadata)
|
||||
console.log('input', input)
|
||||
const DB_URI =
|
||||
env === 'production'
|
||||
? mongoUriProd
|
||||
: env === 'staging'
|
||||
? mongoUriStaging
|
||||
: mongoUriDev
|
||||
|
||||
console.log('DB_URI ', DB_URI)
|
||||
await connectDB(DB_URI)
|
||||
|
||||
const cloudFrontUrl =
|
||||
env === 'production'
|
||||
? cloudFrontUrlProd
|
||||
: env === 'staging'
|
||||
? cloudFrontUrlStaging
|
||||
: cloudFrontUrlDev
|
||||
|
||||
try {
|
||||
await sqs.deleteMessageFromSQS(sqsQueueUrl, receiptHandle)
|
||||
|
||||
const { directoryName } = metadata
|
||||
console.log('directoryName', directoryName)
|
||||
const logPath = `/mnt/efs/potion-voice/${env}/${directoryName}`
|
||||
if (!fs.existsSync(logPath)) {
|
||||
fs.mkdirSync(logPath, { recursive: true })
|
||||
}
|
||||
// update the db model to processing
|
||||
await voiceCloningService.update({ _id, status: 'processing' })
|
||||
await userAudioProfileService.update({
|
||||
_id: userAudioProfileId,
|
||||
status: 'processing',
|
||||
})
|
||||
|
||||
// create directory for userid-useraudioprofileid if not exist
|
||||
const rootPath = `/tmp/${directoryName}`
|
||||
const wavePath = `${rootPath}/wav48/1`
|
||||
if (!fs.existsSync(wavePath)) {
|
||||
fs.mkdirSync(wavePath, { recursive: true })
|
||||
}
|
||||
|
||||
const txtPath = `${rootPath}/txt/1`
|
||||
if (!fs.existsSync(txtPath)) {
|
||||
fs.mkdirSync(txtPath, { recursive: true })
|
||||
}
|
||||
// download the training data files and put it in respective directories
|
||||
for (let index = 0; index < input.length; index++) {
|
||||
const item = input[index]
|
||||
|
||||
const { waveUrl, originalText } = item
|
||||
// download wave file
|
||||
const waveFilePath = `${wavePath}/1_${pad('' + (index + 1))}.wav`
|
||||
|
||||
await getFile(updateUrl(waveUrl, cloudFrontUrl), waveFilePath)
|
||||
|
||||
const txtFilePath = `${txtPath}/1_${pad('' + (index + 1))}.txt`
|
||||
await fs.promises.writeFile(txtFilePath, originalText)
|
||||
}
|
||||
|
||||
const zipFileName = directoryName + '.tgz'
|
||||
|
||||
// /tmp/directoryName.tgz
|
||||
|
||||
await execShellCommand(
|
||||
`cd /tmp && tar czvf ${zipFileName} ${directoryName}`,
|
||||
logPath
|
||||
)
|
||||
console.log('ZIP created ', zipFileName)
|
||||
|
||||
// re-sample audio
|
||||
const SAMPLING_LABEL = `Time Taken for re-sampling ${directoryName}`
|
||||
console.time(SAMPLING_LABEL)
|
||||
|
||||
const outputPath = `/mnt/efs/potion-voice/${env}/${directoryName}`
|
||||
|
||||
const samplingCommand = `python3 ../voice-cloning/prepare_datasets.py --dataset_preset potion_voice_cloning --dataset_archive_path /tmp/${zipFileName} --output_path ${outputPath}`
|
||||
console.log('samplingCommand ', samplingCommand)
|
||||
const samplingResponse = await execShellCommand(
|
||||
samplingCommand,
|
||||
logPath
|
||||
)
|
||||
console.timeEnd(SAMPLING_LABEL)
|
||||
|
||||
// /mnt/efs/potion-voice/${env}/speakrs.pth
|
||||
// /mnt/efs/potion-voice/${env}/txt
|
||||
// /mnt/efs/potion-voice/${env}/${directoryName}/wav
|
||||
|
||||
const outPath = `/mnt/efs/potion-voice/${env}/${directoryName}/sr22050/${directoryName}`
|
||||
|
||||
const resultsPath = outPath + '/results'
|
||||
|
||||
//update pth file for cloning
|
||||
// clone the voice
|
||||
const VOICE_CLONING_LABEL = `Time Taken for voice cloning ${directoryName}`
|
||||
console.time(VOICE_CLONING_LABEL)
|
||||
const trainingModelCommand = `python3 ../voice-cloning/clone_voice.py --baseline_model_path ../voice-cloning/pretrained-models/checkpoint_365000.pth --speaker_dataset_path ${outPath} --speaker_embeddings_path ${
|
||||
outPath + '/speakers.pth'
|
||||
} --output_path ${resultsPath}`
|
||||
|
||||
console.log('Training Model Command', trainingModelCommand)
|
||||
const trainingResponse = await execShellCommand(
|
||||
trainingModelCommand,
|
||||
logPath
|
||||
)
|
||||
|
||||
console.timeEnd(VOICE_CLONING_LABEL)
|
||||
|
||||
let generatedDirectoryName = ''
|
||||
fs.readdirSync(`${resultsPath}/`).forEach((file) => {
|
||||
if (file.includes('vits_potion_clone'))
|
||||
// use output from above to get right path and directory name
|
||||
generatedDirectoryName = file
|
||||
})
|
||||
|
||||
// minimize cloning model
|
||||
const VOICE_MINIMIZE_LABEL = `Time Taken for voice minimizing cloning ${directoryName}`
|
||||
console.time(VOICE_MINIMIZE_LABEL)
|
||||
const minimizeCloningModelCommand = `python3 ../voice-cloning/minimize_cloned_voice_model.py --voice_model_asset_path ${
|
||||
resultsPath + '/' + generatedDirectoryName + '/'
|
||||
} --voice_model_name checkpoint_365200.pth`
|
||||
|
||||
console.log(
|
||||
'Minimize Cloning Model Command',
|
||||
minimizeCloningModelCommand
|
||||
)
|
||||
const minimizeCloning = await execShellCommand(
|
||||
minimizeCloningModelCommand,
|
||||
logPath
|
||||
)
|
||||
console.timeEnd(VOICE_MINIMIZE_LABEL)
|
||||
|
||||
// Add the code to update location of generated model and status into DB
|
||||
await voiceCloningService.update({ _id, status: 'completed' })
|
||||
|
||||
const training_model_path = {
|
||||
voice_model_path: `${resultsPath}/${generatedDirectoryName}/checkpoint_365200.pth`,
|
||||
voice_model_config_path: `${resultsPath}/${generatedDirectoryName}/config.json`,
|
||||
voice_model_speakers_file_path: `${outPath}/speakers.pth`, // TODO update the name to voice model speakers embeddings
|
||||
voice_model_light_path: `${resultsPath}/${generatedDirectoryName}/checkpoint_365200_light.pth`,
|
||||
voice_model_config_light_path: `${resultsPath}/${generatedDirectoryName}/config_light.json`,
|
||||
}
|
||||
|
||||
await userAudioProfileService.update({
|
||||
_id: userAudioProfileId,
|
||||
status: 'completed',
|
||||
training_model_path,
|
||||
})
|
||||
|
||||
// add code to put that model into S3
|
||||
let keys = Object.keys(training_model_path)
|
||||
|
||||
const training_model_s3_path = {}
|
||||
|
||||
for (let index = 0; index < keys.length; index++) {
|
||||
const path = training_model_path[keys[index]]
|
||||
const s3Path = await s3.upload({
|
||||
filePath: path,
|
||||
fileName: `${directoryName}/${path.split('/').pop()}`,
|
||||
bucket: `potion-voice-users-training-model/${env}`,
|
||||
})
|
||||
training_model_s3_path[keys[index]] = s3Path
|
||||
}
|
||||
// add S3 path to user audio profile model
|
||||
await userAudioProfileService.update({
|
||||
_id: userAudioProfileId,
|
||||
training_model_s3_path,
|
||||
})
|
||||
} catch (error) {
|
||||
console.log('error********************', error)
|
||||
Bugsnag.notify(
|
||||
new Error(
|
||||
`Unable to train for voice cloning videos ` + JSON.stringify(job)
|
||||
)
|
||||
)
|
||||
Bugsnag.notify(error)
|
||||
|
||||
// update the db to set status as error
|
||||
await voiceCloningService.update({ _id, status: 'error' })
|
||||
await userAudioProfileService.update({
|
||||
_id: userAudioProfileId,
|
||||
status: 'error',
|
||||
})
|
||||
|
||||
resolve() // to continue working on new jobs
|
||||
}
|
||||
} else {
|
||||
throttleMessageFetching = true
|
||||
}
|
||||
resolve()
|
||||
} catch (error) {
|
||||
console.error('Error while training voice clone', { error })
|
||||
Bugsnag.notify(error)
|
||||
resolve() // to continue working on new jobs
|
||||
} finally {
|
||||
mongoose.connection.close()
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
function sleep(ms) {
|
||||
return new Promise((resolve) => {
|
||||
setTimeout(resolve, ms)
|
||||
})
|
||||
}
|
||||
const init = async () => {
|
||||
console.log('potion Voice Clone Process Started')
|
||||
Bugsnag.start({
|
||||
appVersion: APP_ENV + version,
|
||||
apiKey: process.env.BUGSNAG_BACKEND_KEY,
|
||||
releaseStage: process.env.NODE_ENV,
|
||||
})
|
||||
|
||||
try {
|
||||
while (true) {
|
||||
await processQueue()
|
||||
if (throttleMessageFetching) await sleep(2000)
|
||||
}
|
||||
} catch (error) {
|
||||
Bugsnag.notify(error)
|
||||
}
|
||||
}
|
||||
init()
|
||||
Reference in New Issue
Block a user