333 lines
11 KiB
JavaScript
333 lines
11 KiB
JavaScript
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()
|