Merge pull request #2495 from Risenafis/fix-change-model Fix transformers model changes

db4fe14011ef508d4090ccd52c2a7902c5a55fd7

Cohee <18619528+Cohee1207@users.noreply.github.com>

Signed
1 files changed, +5 -0Showing whitespace changes
src/transformers.mjs+5 -0
@@ -94,8 +94,12 @@ function getModelForTask(task) {
94 */94 */
95async function getPipeline(task, forceModel = '') {95async function getPipeline(task, forceModel = '') {
96 if (tasks[task].pipeline) {96 if (tasks[task].pipeline) {
97 if (forceModel === '' || tasks[task].currentModel === forceModel) {
97 return tasks[task].pipeline;98 return tasks[task].pipeline;
98 }99 }
100 console.log('Disposing transformers.js pipeline for for task', task, 'with model', tasks[task].currentModel);
101 await tasks[task].pipeline.dispose();
102 }
99103
100 const cache_dir = path.join(process.cwd(), 'cache');104 const cache_dir = path.join(process.cwd(), 'cache');
101 const model = forceModel || getModelForTask(task);105 const model = forceModel || getModelForTask(task);
@@ -103,6 +107,7 @@ async function getPipeline(task, forceModel = '') {
103 console.log('Initializing transformers.js pipeline for task', task, 'with model', model);107 console.log('Initializing transformers.js pipeline for task', task, 'with model', model);
104 const instance = await pipeline(task, model, { cache_dir, quantized: tasks[task].quantized ?? true, local_files_only: localOnly });108 const instance = await pipeline(task, model, { cache_dir, quantized: tasks[task].quantized ?? true, local_files_only: localOnly });
105 tasks[task].pipeline = instance;109 tasks[task].pipeline = instance;
110 tasks[task].currentModel = model;
106 return instance;111 return instance;
107}112}
108113