feat: add expressive offline Russian book TTS
This commit is contained in:
@@ -61,3 +61,5 @@
|
|||||||
#*.PDF diff=astextplain
|
#*.PDF diff=astextplain
|
||||||
#*.rtf diff=astextplain
|
#*.rtf diff=astextplain
|
||||||
#*.RTF diff=astextplain
|
#*.RTF diff=astextplain
|
||||||
|
app/src/main/assets/tts/**/*.onnx filter=lfs diff=lfs merge=lfs -text
|
||||||
|
app/src/main/assets/tts/**/dict.tsv filter=lfs diff=lfs merge=lfs -text
|
||||||
|
|||||||
@@ -18,6 +18,9 @@ android {
|
|||||||
vectorDrawables {
|
vectorDrawables {
|
||||||
useSupportLibrary = true
|
useSupportLibrary = true
|
||||||
}
|
}
|
||||||
|
ndk {
|
||||||
|
abiFilters += "arm64-v8a"
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
buildTypes {
|
buildTypes {
|
||||||
@@ -53,6 +56,10 @@ android {
|
|||||||
excludes += "/META-INF/{AL2.0,LGPL2.1}"
|
excludes += "/META-INF/{AL2.0,LGPL2.1}"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
androidResources {
|
||||||
|
noCompress += "onnx"
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
dependencies {
|
dependencies {
|
||||||
@@ -67,6 +74,7 @@ dependencies {
|
|||||||
implementation("androidx.constraintlayout:constraintlayout:2.2.1")
|
implementation("androidx.constraintlayout:constraintlayout:2.2.1")
|
||||||
implementation("androidx.swiperefreshlayout:swiperefreshlayout:1.1.0")
|
implementation("androidx.swiperefreshlayout:swiperefreshlayout:1.1.0")
|
||||||
implementation("androidx.webkit:webkit:1.12.1")
|
implementation("androidx.webkit:webkit:1.12.1")
|
||||||
|
implementation("com.microsoft.onnxruntime:onnxruntime-android:1.27.0")
|
||||||
|
|
||||||
testImplementation("junit:junit:4.13.2")
|
testImplementation("junit:junit:4.13.2")
|
||||||
}
|
}
|
||||||
|
|||||||
Vendored
+4
@@ -5,3 +5,7 @@
|
|||||||
|
|
||||||
# Keep the manifest-declared Argus package installer callback constructable by Android.
|
# Keep the manifest-declared Argus package installer callback constructable by Android.
|
||||||
-keep class xyz.kusoft.argusupdater.ArgusPackageInstallerStatusReceiver { public <init>(); }
|
-keep class xyz.kusoft.argusupdater.ArgusPackageInstallerStatusReceiver { public <init>(); }
|
||||||
|
|
||||||
|
# ONNX Runtime's native JNI layer resolves these Java classes and members by
|
||||||
|
# their original names when converting native OrtValues to Java OnnxValues.
|
||||||
|
-keep class ai.onnxruntime.** { *; }
|
||||||
|
|||||||
@@ -301,6 +301,68 @@
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function visibleTextSlice(element, width, height) {
|
||||||
|
const doc = element && element.ownerDocument;
|
||||||
|
if (!doc || typeof doc.createTreeWalker !== 'function' || typeof doc.createRange !== 'function') return '';
|
||||||
|
const nodeFilter = (doc.defaultView && doc.defaultView.NodeFilter) || global.NodeFilter;
|
||||||
|
if (!nodeFilter) return '';
|
||||||
|
const walker = doc.createTreeWalker(element, nodeFilter.SHOW_TEXT);
|
||||||
|
const range = doc.createRange();
|
||||||
|
let raw = '';
|
||||||
|
let visibleStart = Number.POSITIVE_INFINITY;
|
||||||
|
let visibleEnd = -1;
|
||||||
|
let node;
|
||||||
|
while ((node = walker.nextNode())) {
|
||||||
|
const value = String(node.nodeValue || '');
|
||||||
|
const base = raw.length;
|
||||||
|
raw += value;
|
||||||
|
const words = /\S+/g;
|
||||||
|
let match;
|
||||||
|
while ((match = words.exec(value))) {
|
||||||
|
range.setStart(node, match.index);
|
||||||
|
range.setEnd(node, match.index + match[0].length);
|
||||||
|
const visible = Array.from(range.getClientRects()).some(function (rect) {
|
||||||
|
return rect.width > 0 && rect.height > 0 &&
|
||||||
|
rect.right > 0 && rect.bottom > 0 && rect.left < width && rect.top < height;
|
||||||
|
});
|
||||||
|
if (!visible) continue;
|
||||||
|
visibleStart = Math.min(visibleStart, base + match.index);
|
||||||
|
visibleEnd = Math.max(visibleEnd, base + match.index + match[0].length);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
range.detach();
|
||||||
|
if (!Number.isFinite(visibleStart) || visibleEnd <= visibleStart) return '';
|
||||||
|
|
||||||
|
let start = visibleStart;
|
||||||
|
const before = raw.slice(0, start).replace(/\s+$/g, '');
|
||||||
|
if (before && !/[.!?…][\"'»”)]?$/.test(before)) {
|
||||||
|
const remainder = raw.slice(start, visibleEnd);
|
||||||
|
const nextSentence = /[.!?…]+[\"'»”)]?\s+/.exec(remainder);
|
||||||
|
if (nextSentence) {
|
||||||
|
const candidate = start + nextSentence.index + nextSentence[0].length;
|
||||||
|
if (candidate < visibleEnd) start = candidate;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return raw.slice(start, visibleEnd).replace(/\s+/g, ' ').trim();
|
||||||
|
}
|
||||||
|
|
||||||
|
function rangeSpeechText(range) {
|
||||||
|
if (!range || typeof range.cloneContents !== 'function') return '';
|
||||||
|
const doc = range.startContainer && range.startContainer.ownerDocument;
|
||||||
|
if (!doc) return String(range.toString() || '').trim();
|
||||||
|
const container = doc.createElement('div');
|
||||||
|
container.appendChild(range.cloneContents());
|
||||||
|
const blockSelector = 'p,h1,h2,h3,h4,h5,h6,pre,blockquote,li,td,th';
|
||||||
|
Array.from(container.querySelectorAll(blockSelector)).forEach(function (element) {
|
||||||
|
if (element.querySelector(blockSelector)) return;
|
||||||
|
element.appendChild(doc.createTextNode('\n'));
|
||||||
|
});
|
||||||
|
return String(container.textContent || '')
|
||||||
|
.replace(/[\t\f\v ]+/g, ' ')
|
||||||
|
.replace(/ *\n+ */g, '\n')
|
||||||
|
.trim();
|
||||||
|
}
|
||||||
|
|
||||||
function visibleSpeechText(documents, maxCharacters) {
|
function visibleSpeechText(documents, maxCharacters) {
|
||||||
const limit = clamp(Number(maxCharacters) || 3500, 200, 3500);
|
const limit = clamp(Number(maxCharacters) || 3500, 200, 3500);
|
||||||
const blocks = [];
|
const blocks = [];
|
||||||
@@ -318,13 +380,32 @@
|
|||||||
rect.right > 0 && rect.bottom > 0 && rect.left < width && rect.top < height;
|
rect.right > 0 && rect.bottom > 0 && rect.left < width && rect.top < height;
|
||||||
});
|
});
|
||||||
if (!visible) return;
|
if (!visible) return;
|
||||||
const text = String(element.innerText || element.textContent || '')
|
const text = visibleTextSlice(element, width, height);
|
||||||
.replace(/\s+/g, ' ')
|
|
||||||
.trim();
|
|
||||||
if (!text || seen.has(text)) return;
|
if (!text || seen.has(text)) return;
|
||||||
seen.add(text);
|
seen.add(text);
|
||||||
blocks.push(text);
|
blocks.push(text);
|
||||||
});
|
});
|
||||||
|
if (blocks.length || typeof doc.elementFromPoint !== 'function') return;
|
||||||
|
|
||||||
|
// EPUB pagination can place the current column in a clipped iframe where
|
||||||
|
// getClientRects() does not intersect the document's nominal viewport.
|
||||||
|
// Sampling the actual hit-tested viewport reliably finds the paragraphs
|
||||||
|
// the reader is displaying without reading hidden chapters or controls.
|
||||||
|
const blockSelector = 'p,h1,h2,h3,h4,h5,h6,pre,blockquote,li,td,th';
|
||||||
|
const xSamples = [0.12, 0.5, 0.88];
|
||||||
|
const ySamples = [0.08, 0.22, 0.38, 0.54, 0.70, 0.86, 0.96];
|
||||||
|
ySamples.forEach(function (yRatio) {
|
||||||
|
xSamples.forEach(function (xRatio) {
|
||||||
|
const hit = doc.elementFromPoint(width * xRatio, height * yRatio);
|
||||||
|
const element = hit && hit.closest ? hit.closest(blockSelector) : null;
|
||||||
|
if (!element) return;
|
||||||
|
const text = visibleTextSlice(element, width, height) ||
|
||||||
|
String(element.innerText || element.textContent || '').replace(/\s+/g, ' ').trim();
|
||||||
|
if (!text || seen.has(text)) return;
|
||||||
|
seen.add(text);
|
||||||
|
blocks.push(text);
|
||||||
|
});
|
||||||
|
});
|
||||||
});
|
});
|
||||||
let text = blocks.join('\n').trim();
|
let text = blocks.join('\n').trim();
|
||||||
if (text.length <= limit) return text;
|
if (text.length <= limit) return text;
|
||||||
@@ -363,5 +444,6 @@
|
|||||||
Internal.installReaderFontFaces = installReaderFontFaces;
|
Internal.installReaderFontFaces = installReaderFontFaces;
|
||||||
Internal.waitForPreferredFont = waitForPreferredFont;
|
Internal.waitForPreferredFont = waitForPreferredFont;
|
||||||
Internal.validExternalUrl = validExternalUrl;
|
Internal.validExternalUrl = validExternalUrl;
|
||||||
|
Internal.rangeSpeechText = rangeSpeechText;
|
||||||
Internal.visibleSpeechText = visibleSpeechText;
|
Internal.visibleSpeechText = visibleSpeechText;
|
||||||
})(window);
|
})(window);
|
||||||
|
|||||||
@@ -566,10 +566,95 @@
|
|||||||
const contents = this.rendition && this.rendition.getContents
|
const contents = this.rendition && this.rendition.getContents
|
||||||
? this.rendition.getContents()
|
? this.rendition.getContents()
|
||||||
: [];
|
: [];
|
||||||
|
const location = this.rendition && typeof this.rendition.currentLocation === 'function'
|
||||||
|
? this.rendition.currentLocation()
|
||||||
|
: this.currentLocation;
|
||||||
|
const startCfi = location && location.start && location.start.cfi;
|
||||||
|
const endCfi = location && location.end && location.end.cfi;
|
||||||
|
const limit = Internal.clamp(Number(maxCharacters) || 3500, 200, 3500);
|
||||||
|
if (startCfi && endCfi) {
|
||||||
|
for (let index = 0; index < contents.length; index += 1) {
|
||||||
|
const item = contents[index];
|
||||||
|
if (!item || typeof item.range !== 'function' || !item.document) continue;
|
||||||
|
try {
|
||||||
|
const start = item.range(startCfi);
|
||||||
|
const end = item.range(endCfi);
|
||||||
|
if (!start || !end) continue;
|
||||||
|
const pageRange = item.document.createRange();
|
||||||
|
pageRange.setStart(start.startContainer, start.startOffset);
|
||||||
|
pageRange.setEnd(end.endContainer, end.endOffset);
|
||||||
|
const pageText = Internal.rangeSpeechText(pageRange);
|
||||||
|
let text = pageText;
|
||||||
|
let extendedCharacters = 0;
|
||||||
|
if (text && !/[.!?…][\"'»”)]?\s*$/.test(text)) {
|
||||||
|
const continuationRange = item.document.createRange();
|
||||||
|
continuationRange.setStart(end.endContainer, end.endOffset);
|
||||||
|
continuationRange.selectNodeContents(item.document.body);
|
||||||
|
continuationRange.setStart(end.endContainer, end.endOffset);
|
||||||
|
const continuation = Internal.rangeSpeechText(continuationRange).trimStart();
|
||||||
|
const sentenceEnd = /^[\s\S]*?[.!?…]+[\"'»”)]?/.exec(continuation);
|
||||||
|
if (sentenceEnd) {
|
||||||
|
const extension = sentenceEnd[0].trimEnd();
|
||||||
|
text = (text + ' ' + extension).trim();
|
||||||
|
extendedCharacters = extension.length;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
const startNode = start.startContainer;
|
||||||
|
const startElement = startNode && (startNode.nodeType === 3 ? startNode.parentElement : startNode);
|
||||||
|
const startBlock = startElement && startElement.closest
|
||||||
|
? startElement.closest('p,h1,h2,h3,h4,h5,h6,pre,blockquote,li,td,th')
|
||||||
|
: null;
|
||||||
|
let textBeforeStart = '';
|
||||||
|
if (startBlock) {
|
||||||
|
const beforeRange = item.document.createRange();
|
||||||
|
beforeRange.selectNodeContents(startBlock);
|
||||||
|
beforeRange.setEnd(start.startContainer, start.startOffset);
|
||||||
|
textBeforeStart = String(beforeRange.toString() || '').trimEnd();
|
||||||
|
} else if (startNode && startNode.nodeType === 3 && start.startOffset > 0) {
|
||||||
|
textBeforeStart = String(startNode.nodeValue || '').slice(0, start.startOffset).trimEnd();
|
||||||
|
}
|
||||||
|
const startsInsideSentence = Boolean(textBeforeStart) &&
|
||||||
|
!/[.!?…][\"'»”)]?\s*$/.test(textBeforeStart);
|
||||||
|
if (startsInsideSentence) {
|
||||||
|
const boundary = /[.!?…]+[\"'»”)]?\s+/.exec(text);
|
||||||
|
if (boundary && boundary.index + boundary[0].length < text.length) {
|
||||||
|
text = text.slice(boundary.index + boundary[0].length).trim();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (text.length > limit) {
|
||||||
|
const candidate = text.slice(0, limit + 1);
|
||||||
|
const sentenceEnd = Math.max(
|
||||||
|
candidate.lastIndexOf('. '),
|
||||||
|
candidate.lastIndexOf('! '),
|
||||||
|
candidate.lastIndexOf('? '),
|
||||||
|
candidate.lastIndexOf('… ')
|
||||||
|
);
|
||||||
|
const wordEnd = candidate.lastIndexOf(' ', limit);
|
||||||
|
const cut = sentenceEnd >= Math.floor(limit * .55)
|
||||||
|
? sentenceEnd + 1
|
||||||
|
: (wordEnd > 0 ? wordEnd : limit);
|
||||||
|
text = candidate.slice(0, cut).trim();
|
||||||
|
}
|
||||||
|
if (text) {
|
||||||
|
return {
|
||||||
|
text: text,
|
||||||
|
locator: this.currentLocator(),
|
||||||
|
canAdvance: this.currentProgress < .9999,
|
||||||
|
source: 'epub-cfi-page',
|
||||||
|
pageCharacters: pageText.length,
|
||||||
|
extendedCharacters: extendedCharacters
|
||||||
|
};
|
||||||
|
}
|
||||||
|
} catch (_) {}
|
||||||
|
}
|
||||||
|
}
|
||||||
return {
|
return {
|
||||||
text: Internal.visibleSpeechText(contents.map(function (item) { return item.document; }), maxCharacters),
|
text: Internal.visibleSpeechText(contents.map(function (item) { return item.document; }), maxCharacters),
|
||||||
locator: this.currentLocator(),
|
locator: this.currentLocator(),
|
||||||
canAdvance: this.currentProgress < .9999
|
canAdvance: this.currentProgress < .9999,
|
||||||
|
source: 'visible-fallback',
|
||||||
|
pageCharacters: 0,
|
||||||
|
extendedCharacters: 0
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -222,9 +222,10 @@
|
|||||||
},
|
},
|
||||||
getSpeechPageJson: function (maxCharacters) {
|
getSpeechPageJson: function (maxCharacters) {
|
||||||
try {
|
try {
|
||||||
return JSON.stringify(callEngine('speechPage', [maxCharacters], {
|
const page = state.engine && typeof state.engine.speechPage === 'function'
|
||||||
text: '', locator: null, canAdvance: false
|
? state.engine.speechPage(maxCharacters)
|
||||||
}));
|
: { text: '', locator: null, canAdvance: false };
|
||||||
|
return JSON.stringify(page);
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
Internal.reportError(error, 'api.getSpeechPageJson', true);
|
Internal.reportError(error, 'api.getSpeechPageJson', true);
|
||||||
return JSON.stringify({ text: '', locator: null, canAdvance: false });
|
return JSON.stringify({ text: '', locator: null, canAdvance: false });
|
||||||
|
|||||||
@@ -0,0 +1,201 @@
|
|||||||
|
Apache License
|
||||||
|
Version 2.0, January 2004
|
||||||
|
http://www.apache.org/licenses/
|
||||||
|
|
||||||
|
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||||
|
|
||||||
|
1. Definitions.
|
||||||
|
|
||||||
|
"License" shall mean the terms and conditions for use, reproduction,
|
||||||
|
and distribution as defined by Sections 1 through 9 of this document.
|
||||||
|
|
||||||
|
"Licensor" shall mean the copyright owner or entity authorized by
|
||||||
|
the copyright owner that is granting the License.
|
||||||
|
|
||||||
|
"Legal Entity" shall mean the union of the acting entity and all
|
||||||
|
other entities that control, are controlled by, or are under common
|
||||||
|
control with that entity. For the purposes of this definition,
|
||||||
|
"control" means (i) the power, direct or indirect, to cause the
|
||||||
|
direction or management of such entity, whether by contract or
|
||||||
|
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||||
|
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||||
|
|
||||||
|
"You" (or "Your") shall mean an individual or Legal Entity
|
||||||
|
exercising permissions granted by this License.
|
||||||
|
|
||||||
|
"Source" form shall mean the preferred form for making modifications,
|
||||||
|
including but not limited to software source code, documentation
|
||||||
|
source, and configuration files.
|
||||||
|
|
||||||
|
"Object" form shall mean any form resulting from mechanical
|
||||||
|
transformation or translation of a Source form, including but
|
||||||
|
not limited to compiled object code, generated documentation,
|
||||||
|
and conversions to other media types.
|
||||||
|
|
||||||
|
"Work" shall mean the work of authorship, whether in Source or
|
||||||
|
Object form, made available under the License, as indicated by a
|
||||||
|
copyright notice that is included in or attached to the work
|
||||||
|
(an example is provided in the Appendix below).
|
||||||
|
|
||||||
|
"Derivative Works" shall mean any work, whether in Source or Object
|
||||||
|
form, that is based on (or derived from) the Work and for which the
|
||||||
|
editorial revisions, annotations, elaborations, or other modifications
|
||||||
|
represent, as a whole, an original work of authorship. For the purposes
|
||||||
|
of this License, Derivative Works shall not include works that remain
|
||||||
|
separable from, or merely link (or bind by name) to the interfaces of,
|
||||||
|
the Work and Derivative Works thereof.
|
||||||
|
|
||||||
|
"Contribution" shall mean any work of authorship, including
|
||||||
|
the original version of the Work and any modifications or additions
|
||||||
|
to that Work or Derivative Works thereof, that is intentionally
|
||||||
|
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||||
|
or by an individual or Legal Entity authorized to submit on behalf of
|
||||||
|
the copyright owner. For the purposes of this definition, "submitted"
|
||||||
|
means any form of electronic, verbal, or written communication sent
|
||||||
|
to the Licensor or its representatives, including but not limited to
|
||||||
|
communication on electronic mailing lists, source code control systems,
|
||||||
|
and issue tracking systems that are managed by, or on behalf of, the
|
||||||
|
Licensor for the purpose of discussing and improving the Work, but
|
||||||
|
excluding communication that is conspicuously marked or otherwise
|
||||||
|
designated in writing by the copyright owner as "Not a Contribution."
|
||||||
|
|
||||||
|
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||||
|
on behalf of whom a Contribution has been received by Licensor and
|
||||||
|
subsequently incorporated within the Work.
|
||||||
|
|
||||||
|
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||||
|
this License, each Contributor hereby grants to You a perpetual,
|
||||||
|
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||||
|
copyright license to reproduce, prepare Derivative Works of,
|
||||||
|
publicly display, publicly perform, sublicense, and distribute the
|
||||||
|
Work and such Derivative Works in Source or Object form.
|
||||||
|
|
||||||
|
3. Grant of Patent License. Subject to the terms and conditions of
|
||||||
|
this License, each Contributor hereby grants to You a perpetual,
|
||||||
|
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||||
|
(except as stated in this section) patent license to make, have made,
|
||||||
|
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||||
|
where such license applies only to those patent claims licensable
|
||||||
|
by such Contributor that are necessarily infringed by their
|
||||||
|
Contribution(s) alone or by combination of their Contribution(s)
|
||||||
|
with the Work to which such Contribution(s) was submitted. If You
|
||||||
|
institute patent litigation against any entity (including a
|
||||||
|
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||||
|
or a Contribution incorporated within the Work constitutes direct
|
||||||
|
or contributory patent infringement, then any patent licenses
|
||||||
|
granted to You under this License for that Work shall terminate
|
||||||
|
as of the date such litigation is filed.
|
||||||
|
|
||||||
|
4. Redistribution. You may reproduce and distribute copies of the
|
||||||
|
Work or Derivative Works thereof in any medium, with or without
|
||||||
|
modifications, and in Source or Object form, provided that You
|
||||||
|
meet the following conditions:
|
||||||
|
|
||||||
|
(a) You must give any other recipients of the Work or
|
||||||
|
Derivative Works a copy of this License; and
|
||||||
|
|
||||||
|
(b) You must cause any modified files to carry prominent notices
|
||||||
|
stating that You changed the files; and
|
||||||
|
|
||||||
|
(c) You must retain, in the Source form of any Derivative Works
|
||||||
|
that You distribute, all copyright, patent, trademark, and
|
||||||
|
attribution notices from the Source form of the Work,
|
||||||
|
excluding those notices that do not pertain to any part of
|
||||||
|
the Derivative Works; and
|
||||||
|
|
||||||
|
(d) If the Work includes a "NOTICE" text file as part of its
|
||||||
|
distribution, then any Derivative Works that You distribute must
|
||||||
|
include a readable copy of the attribution notices contained
|
||||||
|
within such NOTICE file, excluding those notices that do not
|
||||||
|
pertain to any part of the Derivative Works, in at least one
|
||||||
|
of the following places: within a NOTICE text file distributed
|
||||||
|
as part of the Derivative Works; within the Source form or
|
||||||
|
documentation, if provided along with the Derivative Works; or,
|
||||||
|
within a display generated by the Derivative Works, if and
|
||||||
|
wherever such third-party notices normally appear. The contents
|
||||||
|
of the NOTICE file are for informational purposes only and
|
||||||
|
do not modify the License. You may add Your own attribution
|
||||||
|
notices within Derivative Works that You distribute, alongside
|
||||||
|
or as an addendum to the NOTICE text from the Work, provided
|
||||||
|
that such additional attribution notices cannot be construed
|
||||||
|
as modifying the License.
|
||||||
|
|
||||||
|
You may add Your own copyright statement to Your modifications and
|
||||||
|
may provide additional or different license terms and conditions
|
||||||
|
for use, reproduction, or distribution of Your modifications, or
|
||||||
|
for any such Derivative Works as a whole, provided Your use,
|
||||||
|
reproduction, and distribution of the Work otherwise complies with
|
||||||
|
the conditions stated in this License.
|
||||||
|
|
||||||
|
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||||
|
any Contribution intentionally submitted for inclusion in the Work
|
||||||
|
by You to the Licensor shall be under the terms and conditions of
|
||||||
|
this License, without any additional terms or conditions.
|
||||||
|
Notwithstanding the above, nothing herein shall supersede or modify
|
||||||
|
the terms of any separate license agreement you may have executed
|
||||||
|
with Licensor regarding such Contributions.
|
||||||
|
|
||||||
|
6. Trademarks. This License does not grant permission to use the trade
|
||||||
|
names, trademarks, service marks, or product names of the Licensor,
|
||||||
|
except as required for reasonable and customary use in describing the
|
||||||
|
origin of the Work and reproducing the content of the NOTICE file.
|
||||||
|
|
||||||
|
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||||
|
agreed to in writing, Licensor provides the Work (and each
|
||||||
|
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||||
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||||
|
implied, including, without limitation, any warranties or conditions
|
||||||
|
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||||
|
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||||
|
appropriateness of using or redistributing the Work and assume any
|
||||||
|
risks associated with Your exercise of permissions under this License.
|
||||||
|
|
||||||
|
8. Limitation of Liability. In no event and under no legal theory,
|
||||||
|
whether in tort (including negligence), contract, or otherwise,
|
||||||
|
unless required by applicable law (such as deliberate and grossly
|
||||||
|
negligent acts) or agreed to in writing, shall any Contributor be
|
||||||
|
liable to You for damages, including any direct, indirect, special,
|
||||||
|
incidental, or consequential damages of any character arising as a
|
||||||
|
result of this License or out of the use or inability to use the
|
||||||
|
Work (including but not limited to damages for loss of goodwill,
|
||||||
|
work stoppage, computer failure or malfunction, or any and all
|
||||||
|
other commercial damages or losses), even if such Contributor
|
||||||
|
has been advised of the possibility of such damages.
|
||||||
|
|
||||||
|
9. Accepting Warranty or Additional Liability. While redistributing
|
||||||
|
the Work or Derivative Works thereof, You may choose to offer,
|
||||||
|
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||||
|
or other liability obligations and/or rights consistent with this
|
||||||
|
License. However, in accepting such obligations, You may act only
|
||||||
|
on Your own behalf and on Your sole responsibility, not on behalf
|
||||||
|
of any other Contributor, and only if You agree to indemnify,
|
||||||
|
defend, and hold each Contributor harmless for any liability
|
||||||
|
incurred by, or claims asserted against, such Contributor by reason
|
||||||
|
of your accepting any such warranty or additional liability.
|
||||||
|
|
||||||
|
END OF TERMS AND CONDITIONS
|
||||||
|
|
||||||
|
APPENDIX: How to apply the Apache License to your work.
|
||||||
|
|
||||||
|
To apply the Apache License to your work, attach the following
|
||||||
|
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||||
|
replaced with your own identifying information. (Don't include
|
||||||
|
the brackets!) The text should be enclosed in the appropriate
|
||||||
|
comment syntax for the file format. We also recommend that a
|
||||||
|
file or class name and description of purpose be included on the
|
||||||
|
same "printed page" as the copyright notice for easier
|
||||||
|
identification within third-party archives.
|
||||||
|
|
||||||
|
Copyright [yyyy] [name of copyright owner]
|
||||||
|
|
||||||
|
Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
you may not use this file except in compliance with the License.
|
||||||
|
You may obtain a copy of the License at
|
||||||
|
|
||||||
|
http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
|
||||||
|
Unless required by applicable law or agreed to in writing, software
|
||||||
|
distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
See the License for the specific language governing permissions and
|
||||||
|
limitations under the License.
|
||||||
Binary file not shown.
Binary file not shown.
|
File diff suppressed because it is too large
Load Diff
Binary file not shown.
File diff suppressed because it is too large
Load Diff
@@ -520,7 +520,7 @@ class ReaderActivity : AppCompatActivity() {
|
|||||||
binding.readerVoiceButton.contentDescription = if (active) {
|
binding.readerVoiceButton.contentDescription = if (active) {
|
||||||
"Приостановить озвучивание"
|
"Приостановить озвучивание"
|
||||||
} else {
|
} else {
|
||||||
"Слушать книгу. Удерживайте для настройки скорости и интонации"
|
"Слушать книгу. Удерживайте для настройки голоса, темпа и пауз"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -401,7 +401,10 @@ class ReaderWebController(
|
|||||||
data class ReaderSpeechPage(
|
data class ReaderSpeechPage(
|
||||||
val text: String,
|
val text: String,
|
||||||
val locator: String?,
|
val locator: String?,
|
||||||
val canAdvance: Boolean
|
val canAdvance: Boolean,
|
||||||
|
val source: String,
|
||||||
|
val pageCharacters: Int,
|
||||||
|
val extendedCharacters: Int
|
||||||
) {
|
) {
|
||||||
companion object {
|
companion object {
|
||||||
internal fun fromJson(json: JSONObject): ReaderSpeechPage? {
|
internal fun fromJson(json: JSONObject): ReaderSpeechPage? {
|
||||||
@@ -415,7 +418,10 @@ data class ReaderSpeechPage(
|
|||||||
return ReaderSpeechPage(
|
return ReaderSpeechPage(
|
||||||
text = text,
|
text = text,
|
||||||
locator = locator,
|
locator = locator,
|
||||||
canAdvance = json.optBoolean("canAdvance", true)
|
canAdvance = json.optBoolean("canAdvance", true),
|
||||||
|
source = json.optString("source", "unknown"),
|
||||||
|
pageCharacters = json.optInt("pageCharacters", text.length),
|
||||||
|
extendedCharacters = json.optInt("extendedCharacters", 0)
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,25 +1,15 @@
|
|||||||
package com.aletheia.app.ui.reader.tts
|
package com.aletheia.app.ui.reader.tts
|
||||||
|
|
||||||
import android.content.Context
|
import android.content.Context
|
||||||
import android.media.AudioAttributes
|
import android.util.Log
|
||||||
import android.speech.tts.TextToSpeech
|
|
||||||
import android.speech.tts.UtteranceProgressListener
|
|
||||||
import com.aletheia.app.ui.reader.ReaderSpeechPage
|
import com.aletheia.app.ui.reader.ReaderSpeechPage
|
||||||
import com.aletheia.app.ui.reader.ReaderWebController
|
import com.aletheia.app.ui.reader.ReaderWebController
|
||||||
import java.util.Locale
|
import java.util.concurrent.CancellationException
|
||||||
import java.util.UUID
|
|
||||||
import kotlinx.coroutines.CoroutineScope
|
import kotlinx.coroutines.CoroutineScope
|
||||||
import kotlinx.coroutines.Job
|
import kotlinx.coroutines.Job
|
||||||
import kotlinx.coroutines.delay
|
import kotlinx.coroutines.delay
|
||||||
import kotlinx.coroutines.launch
|
import kotlinx.coroutines.launch
|
||||||
|
|
||||||
/**
|
|
||||||
* Reads book pages through an installed Russian voice that explicitly declares
|
|
||||||
* that it does not require a network connection.
|
|
||||||
*
|
|
||||||
* The controller deliberately refuses network-only voices. A project-owned
|
|
||||||
* neural model can replace this engine behind the same reader-page contract.
|
|
||||||
*/
|
|
||||||
class ReaderSpeechController(
|
class ReaderSpeechController(
|
||||||
context: Context,
|
context: Context,
|
||||||
private val scope: CoroutineScope,
|
private val scope: CoroutineScope,
|
||||||
@@ -31,10 +21,11 @@ class ReaderSpeechController(
|
|||||||
|
|
||||||
private val appContext = context.applicationContext
|
private val appContext = context.applicationContext
|
||||||
private val preferences = appContext.getSharedPreferences(PREFERENCES_NAME, Context.MODE_PRIVATE)
|
private val preferences = appContext.getSharedPreferences(PREFERENCES_NAME, Context.MODE_PRIVATE)
|
||||||
private var tts: TextToSpeech? = null
|
private val engine = VoskSpeechEngine(appContext)
|
||||||
private var state = State.IDLE
|
private var state = State.IDLE
|
||||||
private var initializationComplete = false
|
private var initializationComplete = false
|
||||||
private var resumeAfterInitialization = false
|
private var resumeAfterInitialization = false
|
||||||
|
private var initializationJob: Job? = null
|
||||||
private var playbackJob: Job? = null
|
private var playbackJob: Job? = null
|
||||||
private var currentPage: ReaderSpeechPage? = null
|
private var currentPage: ReaderSpeechPage? = null
|
||||||
private var lastSpokenSignature: String? = null
|
private var lastSpokenSignature: String? = null
|
||||||
@@ -45,6 +36,16 @@ class ReaderSpeechController(
|
|||||||
val pitch: Float
|
val pitch: Float
|
||||||
get() = preferences.getFloat(KEY_PITCH, DEFAULT_PITCH).coerceIn(MIN_PITCH, MAX_PITCH)
|
get() = preferences.getFloat(KEY_PITCH, DEFAULT_PITCH).coerceIn(MIN_PITCH, MAX_PITCH)
|
||||||
|
|
||||||
|
val pauseScale: Float
|
||||||
|
get() = preferences.getFloat(KEY_PAUSE_SCALE, DEFAULT_PAUSE_SCALE).coerceIn(MIN_PAUSE_SCALE, MAX_PAUSE_SCALE)
|
||||||
|
|
||||||
|
val articulationScale: Float
|
||||||
|
get() = preferences.getFloat(KEY_ARTICULATION_SCALE, DEFAULT_ARTICULATION_SCALE)
|
||||||
|
.coerceIn(MIN_ARTICULATION_SCALE, MAX_ARTICULATION_SCALE)
|
||||||
|
|
||||||
|
val speakerId: Int
|
||||||
|
get() = preferences.getInt(KEY_SPEAKER_ID, DEFAULT_SPEAKER_ID).coerceIn(0, VOICES.lastIndex)
|
||||||
|
|
||||||
fun toggle() {
|
fun toggle() {
|
||||||
when (state) {
|
when (state) {
|
||||||
State.PLAYING, State.INITIALIZING -> pause()
|
State.PLAYING, State.INITIALIZING -> pause()
|
||||||
@@ -58,16 +59,17 @@ class ReaderSpeechController(
|
|||||||
initialize()
|
initialize()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
val resumePausedPage = state == State.PAUSED
|
currentPage = null
|
||||||
|
lastSpokenSignature = null
|
||||||
updateState(State.PLAYING)
|
updateState(State.PLAYING)
|
||||||
speakCurrentPage(reusePausedPage = resumePausedPage)
|
speakCurrentPage(reusePausedPage = false)
|
||||||
}
|
}
|
||||||
|
|
||||||
fun pause() {
|
fun pause() {
|
||||||
resumeAfterInitialization = false
|
resumeAfterInitialization = false
|
||||||
playbackJob?.cancel()
|
playbackJob?.cancel()
|
||||||
playbackJob = null
|
playbackJob = null
|
||||||
tts?.stop()
|
engine.stop()
|
||||||
updateState(State.PAUSED)
|
updateState(State.PAUSED)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -75,94 +77,61 @@ class ReaderSpeechController(
|
|||||||
resumeAfterInitialization = false
|
resumeAfterInitialization = false
|
||||||
playbackJob?.cancel()
|
playbackJob?.cancel()
|
||||||
playbackJob = null
|
playbackJob = null
|
||||||
tts?.stop()
|
engine.stop()
|
||||||
currentPage = null
|
currentPage = null
|
||||||
lastSpokenSignature = null
|
lastSpokenSignature = null
|
||||||
updateState(State.IDLE)
|
updateState(State.IDLE)
|
||||||
}
|
}
|
||||||
|
|
||||||
fun updateSettings(rate: Float, pitch: Float) {
|
fun updateSettings(
|
||||||
|
rate: Float,
|
||||||
|
pitch: Float,
|
||||||
|
pauseScale: Float,
|
||||||
|
articulationScale: Float,
|
||||||
|
speakerId: Int
|
||||||
|
) {
|
||||||
preferences.edit()
|
preferences.edit()
|
||||||
.putFloat(KEY_RATE, rate.coerceIn(MIN_RATE, MAX_RATE))
|
.putFloat(KEY_RATE, rate.coerceIn(MIN_RATE, MAX_RATE))
|
||||||
.putFloat(KEY_PITCH, pitch.coerceIn(MIN_PITCH, MAX_PITCH))
|
.putFloat(KEY_PITCH, pitch.coerceIn(MIN_PITCH, MAX_PITCH))
|
||||||
|
.putFloat(KEY_PAUSE_SCALE, pauseScale.coerceIn(MIN_PAUSE_SCALE, MAX_PAUSE_SCALE))
|
||||||
|
.putFloat(
|
||||||
|
KEY_ARTICULATION_SCALE,
|
||||||
|
articulationScale.coerceIn(MIN_ARTICULATION_SCALE, MAX_ARTICULATION_SCALE)
|
||||||
|
)
|
||||||
|
.putInt(KEY_SPEAKER_ID, speakerId.coerceIn(0, VOICES.lastIndex))
|
||||||
.apply()
|
.apply()
|
||||||
applySpeechSettings()
|
|
||||||
if (state == State.PLAYING) {
|
if (state == State.PLAYING) {
|
||||||
tts?.stop()
|
engine.stop()
|
||||||
speakCurrentPage(reusePausedPage = true)
|
speakCurrentPage(reusePausedPage = true)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fun destroy() {
|
fun destroy() {
|
||||||
|
initializationJob?.cancel()
|
||||||
playbackJob?.cancel()
|
playbackJob?.cancel()
|
||||||
|
initializationJob = null
|
||||||
playbackJob = null
|
playbackJob = null
|
||||||
tts?.stop()
|
engine.destroy()
|
||||||
tts?.shutdown()
|
|
||||||
tts = null
|
|
||||||
initializationComplete = false
|
initializationComplete = false
|
||||||
}
|
}
|
||||||
|
|
||||||
private fun initialize() {
|
private fun initialize() {
|
||||||
if (tts != null) {
|
if (initializationJob?.isActive == true) return
|
||||||
updateState(State.INITIALIZING)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
updateState(State.INITIALIZING)
|
updateState(State.INITIALIZING)
|
||||||
tts = TextToSpeech(appContext) { status ->
|
initializationJob = scope.launch {
|
||||||
if (status != TextToSpeech.SUCCESS) {
|
try {
|
||||||
fail("Не удалось запустить синтез речи на устройстве")
|
engine.initialize()
|
||||||
return@TextToSpeech
|
initializationComplete = true
|
||||||
}
|
if (resumeAfterInitialization) {
|
||||||
val engine = tts ?: return@TextToSpeech
|
resumeAfterInitialization = false
|
||||||
val russianLocale = Locale.Builder().setLanguage("ru").setRegion("RU").build()
|
updateState(State.PLAYING)
|
||||||
if (engine.setLanguage(russianLocale) < TextToSpeech.LANG_AVAILABLE) {
|
speakCurrentPage(reusePausedPage = false)
|
||||||
fail("Русский язык не поддерживается установленным синтезатором речи")
|
} else {
|
||||||
return@TextToSpeech
|
updateState(State.IDLE)
|
||||||
}
|
|
||||||
val offlineRussianVoice = engine.voices
|
|
||||||
?.asSequence()
|
|
||||||
?.filter { it.locale.language.equals("ru", ignoreCase = true) }
|
|
||||||
?.filterNot { it.isNetworkConnectionRequired }
|
|
||||||
?.sortedWith(compareByDescending<android.speech.tts.Voice> { it.quality }.thenBy { it.name })
|
|
||||||
?.firstOrNull()
|
|
||||||
if (offlineRussianVoice == null) {
|
|
||||||
fail("На телефоне не установлен русский офлайн-голос")
|
|
||||||
return@TextToSpeech
|
|
||||||
}
|
|
||||||
if (engine.setVoice(offlineRussianVoice) == TextToSpeech.ERROR) {
|
|
||||||
fail("Не удалось выбрать русский офлайн-голос")
|
|
||||||
return@TextToSpeech
|
|
||||||
}
|
|
||||||
engine.setAudioAttributes(
|
|
||||||
AudioAttributes.Builder()
|
|
||||||
.setUsage(AudioAttributes.USAGE_MEDIA)
|
|
||||||
.setContentType(AudioAttributes.CONTENT_TYPE_SPEECH)
|
|
||||||
.build()
|
|
||||||
)
|
|
||||||
engine.setOnUtteranceProgressListener(object : UtteranceProgressListener() {
|
|
||||||
override fun onStart(utteranceId: String?) = Unit
|
|
||||||
|
|
||||||
override fun onDone(utteranceId: String?) {
|
|
||||||
scope.launch { advanceAfterPage() }
|
|
||||||
}
|
}
|
||||||
|
} catch (error: Throwable) {
|
||||||
@Deprecated("Deprecated in Java")
|
Log.e(TAG, "Offline voice initialization failed", error)
|
||||||
override fun onError(utteranceId: String?) {
|
fail("Не удалось запустить встроенный офлайн-голос Vosk: ${error.message ?: error.javaClass.simpleName}")
|
||||||
scope.launch { fail("Ошибка синтеза речи") }
|
|
||||||
}
|
|
||||||
|
|
||||||
override fun onError(utteranceId: String?, errorCode: Int) {
|
|
||||||
scope.launch { fail("Ошибка синтеза речи: $errorCode") }
|
|
||||||
}
|
|
||||||
})
|
|
||||||
initializationComplete = true
|
|
||||||
applySpeechSettings()
|
|
||||||
if (resumeAfterInitialization) {
|
|
||||||
resumeAfterInitialization = false
|
|
||||||
updateState(State.PLAYING)
|
|
||||||
speakCurrentPage(reusePausedPage = false)
|
|
||||||
} else {
|
|
||||||
updateState(State.IDLE)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -170,29 +139,37 @@ class ReaderSpeechController(
|
|||||||
private fun speakCurrentPage(reusePausedPage: Boolean) {
|
private fun speakCurrentPage(reusePausedPage: Boolean) {
|
||||||
playbackJob?.cancel()
|
playbackJob?.cancel()
|
||||||
playbackJob = scope.launch {
|
playbackJob = scope.launch {
|
||||||
val page = if (reusePausedPage) currentPage else null
|
try {
|
||||||
?: reader.getSpeechPage(TextToSpeech.getMaxSpeechInputLength().coerceAtMost(MAX_PAGE_CHARS))
|
val page = if (reusePausedPage) currentPage else null
|
||||||
if (state != State.PLAYING) return@launch
|
?: reader.getSpeechPage(MAX_PAGE_CHARS)
|
||||||
if (page == null || page.text.isBlank()) {
|
if (state != State.PLAYING) return@launch
|
||||||
fail("На текущей странице нет текста для озвучивания")
|
if (page == null || page.text.isBlank()) {
|
||||||
return@launch
|
fail("На текущей странице нет текста для озвучивания")
|
||||||
|
return@launch
|
||||||
|
}
|
||||||
|
Log.d(
|
||||||
|
TAG,
|
||||||
|
"Speech page extracted: chars=${page.text.length}, pageChars=${page.pageCharacters}, " +
|
||||||
|
"extended=${page.extendedCharacters}, canAdvance=${page.canAdvance}, source=${page.source}"
|
||||||
|
)
|
||||||
|
currentPage = page
|
||||||
|
val signature = page.signature()
|
||||||
|
if (!reusePausedPage && signature == lastSpokenSignature) {
|
||||||
|
stop()
|
||||||
|
onMessage("Достигнут конец книги")
|
||||||
|
return@launch
|
||||||
|
}
|
||||||
|
lastSpokenSignature = signature
|
||||||
|
engine.speak(page.text, rate, pitch, pauseScale, articulationScale)
|
||||||
|
advanceAfterPage()
|
||||||
|
} catch (_: CancellationException) {
|
||||||
|
// Pause and stop cancel the active synthesis/playback job.
|
||||||
|
} catch (error: Throwable) {
|
||||||
|
if (state == State.PLAYING) {
|
||||||
|
Log.e(TAG, "Offline speech playback failed", error)
|
||||||
|
fail("Ошибка встроенного синтеза речи: ${error.message ?: error.javaClass.simpleName}")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
currentPage = page
|
|
||||||
val signature = page.signature()
|
|
||||||
if (!reusePausedPage && signature == lastSpokenSignature) {
|
|
||||||
stop()
|
|
||||||
onMessage("Достигнут конец книги")
|
|
||||||
return@launch
|
|
||||||
}
|
|
||||||
lastSpokenSignature = signature
|
|
||||||
applySpeechSettings()
|
|
||||||
val result = tts?.speak(
|
|
||||||
page.text,
|
|
||||||
TextToSpeech.QUEUE_FLUSH,
|
|
||||||
null,
|
|
||||||
UUID.randomUUID().toString()
|
|
||||||
) ?: TextToSpeech.ERROR
|
|
||||||
if (result == TextToSpeech.ERROR) fail("Не удалось начать озвучивание")
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -210,20 +187,16 @@ class ReaderSpeechController(
|
|||||||
if (state == State.PLAYING) speakCurrentPage(reusePausedPage = false)
|
if (state == State.PLAYING) speakCurrentPage(reusePausedPage = false)
|
||||||
}
|
}
|
||||||
|
|
||||||
private fun applySpeechSettings() {
|
|
||||||
tts?.setSpeechRate(rate)
|
|
||||||
tts?.setPitch(pitch)
|
|
||||||
}
|
|
||||||
|
|
||||||
private fun fail(message: String) {
|
private fun fail(message: String) {
|
||||||
playbackJob?.cancel()
|
playbackJob?.cancel()
|
||||||
playbackJob = null
|
playbackJob = null
|
||||||
tts?.stop()
|
engine.stop()
|
||||||
updateState(State.ERROR)
|
updateState(State.ERROR)
|
||||||
onMessage(message)
|
onMessage(message)
|
||||||
}
|
}
|
||||||
|
|
||||||
private fun updateState(next: State) {
|
private fun updateState(next: State) {
|
||||||
|
Log.d(TAG, "State: $state -> $next")
|
||||||
state = next
|
state = next
|
||||||
onStateChanged(next)
|
onStateChanged(next)
|
||||||
}
|
}
|
||||||
@@ -231,15 +204,27 @@ class ReaderSpeechController(
|
|||||||
private fun ReaderSpeechPage.signature(): String = "$locator|${text.take(160)}"
|
private fun ReaderSpeechPage.signature(): String = "$locator|${text.take(160)}"
|
||||||
|
|
||||||
companion object {
|
companion object {
|
||||||
|
private const val TAG = "AletheiaVoskTts"
|
||||||
const val MIN_RATE = 0.55f
|
const val MIN_RATE = 0.55f
|
||||||
const val MAX_RATE = 1.80f
|
const val MAX_RATE = 1.80f
|
||||||
const val MIN_PITCH = 0.75f
|
const val MIN_PITCH = 0.75f
|
||||||
const val MAX_PITCH = 1.30f
|
const val MAX_PITCH = 1.30f
|
||||||
|
const val MIN_PAUSE_SCALE = 0.50f
|
||||||
|
const val MAX_PAUSE_SCALE = 2.00f
|
||||||
|
const val MIN_ARTICULATION_SCALE = 0.90f
|
||||||
|
const val MAX_ARTICULATION_SCALE = 1.20f
|
||||||
|
val VOICES = listOf("Мужской 1")
|
||||||
private const val DEFAULT_RATE = 1.0f
|
private const val DEFAULT_RATE = 1.0f
|
||||||
private const val DEFAULT_PITCH = 1.0f
|
private const val DEFAULT_PITCH = 1.0f
|
||||||
|
private const val DEFAULT_PAUSE_SCALE = 1.0f
|
||||||
|
private const val DEFAULT_ARTICULATION_SCALE = 1.06f
|
||||||
|
private const val DEFAULT_SPEAKER_ID = 2
|
||||||
private const val PREFERENCES_NAME = "reader_speech"
|
private const val PREFERENCES_NAME = "reader_speech"
|
||||||
private const val KEY_RATE = "rate"
|
private const val KEY_RATE = "rate"
|
||||||
private const val KEY_PITCH = "pitch"
|
private const val KEY_PITCH = "pitch"
|
||||||
|
private const val KEY_PAUSE_SCALE = "pause_scale"
|
||||||
|
private const val KEY_ARTICULATION_SCALE = "articulation_scale"
|
||||||
|
private const val KEY_SPEAKER_ID = "speaker_id"
|
||||||
private const val MAX_PAGE_CHARS = 3_500
|
private const val MAX_PAGE_CHARS = 3_500
|
||||||
private const val PAGE_TURN_DELAY_MS = 450L
|
private const val PAGE_TURN_DELAY_MS = 450L
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,8 +2,10 @@ package com.aletheia.app.ui.reader.tts
|
|||||||
|
|
||||||
import android.content.Context
|
import android.content.Context
|
||||||
import android.view.ViewGroup
|
import android.view.ViewGroup
|
||||||
|
import android.widget.ArrayAdapter
|
||||||
import android.widget.LinearLayout
|
import android.widget.LinearLayout
|
||||||
import android.widget.SeekBar
|
import android.widget.SeekBar
|
||||||
|
import android.widget.Spinner
|
||||||
import android.widget.TextView
|
import android.widget.TextView
|
||||||
import com.google.android.material.dialog.MaterialAlertDialogBuilder
|
import com.google.android.material.dialog.MaterialAlertDialogBuilder
|
||||||
import java.util.Locale
|
import java.util.Locale
|
||||||
@@ -18,10 +20,25 @@ object ReaderSpeechSettingsDialog {
|
|||||||
val vertical = (8 * density).roundToInt()
|
val vertical = (8 * density).roundToInt()
|
||||||
setPadding(horizontal, vertical, horizontal, vertical)
|
setPadding(horizontal, vertical, horizontal, vertical)
|
||||||
}
|
}
|
||||||
|
val voiceTitle = TextView(context).apply {
|
||||||
|
text = "Голос Vosk TTS 0.9"
|
||||||
|
textSize = 15f
|
||||||
|
setPadding(0, 16, 0, 4)
|
||||||
|
}
|
||||||
|
val voice = Spinner(context).apply {
|
||||||
|
adapter = ArrayAdapter(
|
||||||
|
context,
|
||||||
|
android.R.layout.simple_spinner_dropdown_item,
|
||||||
|
ReaderSpeechController.VOICES
|
||||||
|
)
|
||||||
|
setSelection(controller.speakerId)
|
||||||
|
}
|
||||||
|
content.addView(voiceTitle)
|
||||||
|
content.addView(voice)
|
||||||
val rateValue = TextView(context)
|
val rateValue = TextView(context)
|
||||||
val rate = slider(
|
val rate = slider(
|
||||||
context = context,
|
context = context,
|
||||||
label = "Скорость",
|
label = "Темп чтения",
|
||||||
value = controller.rate,
|
value = controller.rate,
|
||||||
min = ReaderSpeechController.MIN_RATE,
|
min = ReaderSpeechController.MIN_RATE,
|
||||||
max = ReaderSpeechController.MAX_RATE,
|
max = ReaderSpeechController.MAX_RATE,
|
||||||
@@ -31,23 +48,54 @@ object ReaderSpeechSettingsDialog {
|
|||||||
val pitchValue = TextView(context)
|
val pitchValue = TextView(context)
|
||||||
val pitch = slider(
|
val pitch = slider(
|
||||||
context = context,
|
context = context,
|
||||||
label = "Интонация",
|
label = "Высота голоса",
|
||||||
value = controller.pitch,
|
value = controller.pitch,
|
||||||
min = ReaderSpeechController.MIN_PITCH,
|
min = ReaderSpeechController.MIN_PITCH,
|
||||||
max = ReaderSpeechController.MAX_PITCH,
|
max = ReaderSpeechController.MAX_PITCH,
|
||||||
valueView = pitchValue,
|
valueView = pitchValue,
|
||||||
parent = content
|
parent = content
|
||||||
)
|
)
|
||||||
|
val pauseValue = TextView(context)
|
||||||
|
val pause = slider(
|
||||||
|
context = context,
|
||||||
|
label = "Паузы на знаках препинания",
|
||||||
|
value = controller.pauseScale,
|
||||||
|
min = ReaderSpeechController.MIN_PAUSE_SCALE,
|
||||||
|
max = ReaderSpeechController.MAX_PAUSE_SCALE,
|
||||||
|
valueView = pauseValue,
|
||||||
|
parent = content
|
||||||
|
)
|
||||||
|
val articulationValue = TextView(context)
|
||||||
|
val articulation = slider(
|
||||||
|
context = context,
|
||||||
|
label = "Разборчивость слов и букв",
|
||||||
|
value = controller.articulationScale,
|
||||||
|
min = ReaderSpeechController.MIN_ARTICULATION_SCALE,
|
||||||
|
max = ReaderSpeechController.MAX_ARTICULATION_SCALE,
|
||||||
|
valueView = articulationValue,
|
||||||
|
parent = content
|
||||||
|
)
|
||||||
|
|
||||||
MaterialAlertDialogBuilder(context)
|
MaterialAlertDialogBuilder(context)
|
||||||
.setTitle("Озвучивание книги")
|
.setTitle("Озвучивание книги")
|
||||||
.setMessage("Голос работает офлайн. Интонация регулирует высоту голоса.")
|
.setMessage("Все параметры применяются к встроенному офлайн-голосу. Настройки меняются плавно.")
|
||||||
.setView(content)
|
.setView(content)
|
||||||
.setNegativeButton("Отмена", null)
|
.setNegativeButton("Отмена", null)
|
||||||
.setPositiveButton("Сохранить") { _, _ ->
|
.setPositiveButton("Сохранить") { _, _ ->
|
||||||
controller.updateSettings(
|
controller.updateSettings(
|
||||||
progressToValue(rate.progress, ReaderSpeechController.MIN_RATE, ReaderSpeechController.MAX_RATE),
|
progressToValue(rate.progress, ReaderSpeechController.MIN_RATE, ReaderSpeechController.MAX_RATE),
|
||||||
progressToValue(pitch.progress, ReaderSpeechController.MIN_PITCH, ReaderSpeechController.MAX_PITCH)
|
progressToValue(pitch.progress, ReaderSpeechController.MIN_PITCH, ReaderSpeechController.MAX_PITCH),
|
||||||
|
progressToValue(
|
||||||
|
pause.progress,
|
||||||
|
ReaderSpeechController.MIN_PAUSE_SCALE,
|
||||||
|
ReaderSpeechController.MAX_PAUSE_SCALE
|
||||||
|
),
|
||||||
|
progressToValue(
|
||||||
|
articulation.progress,
|
||||||
|
ReaderSpeechController.MIN_ARTICULATION_SCALE,
|
||||||
|
ReaderSpeechController.MAX_ARTICULATION_SCALE
|
||||||
|
),
|
||||||
|
voice.selectedItemPosition
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
.show()
|
.show()
|
||||||
@@ -97,5 +145,5 @@ object ReaderSpeechSettingsDialog {
|
|||||||
private fun progressToValue(progress: Int, min: Float, max: Float): Float =
|
private fun progressToValue(progress: Int, min: Float, max: Float): Float =
|
||||||
min + (max - min) * progress.coerceIn(0, SLIDER_STEPS) / SLIDER_STEPS
|
min + (max - min) * progress.coerceIn(0, SLIDER_STEPS) / SLIDER_STEPS
|
||||||
|
|
||||||
private const val SLIDER_STEPS = 100
|
private const val SLIDER_STEPS = 300
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,141 @@
|
|||||||
|
package com.aletheia.app.ui.reader.tts
|
||||||
|
|
||||||
|
internal object RussianNumberNormalizer {
|
||||||
|
private val yearInPrepositionalPattern = Regex(
|
||||||
|
"(?<![\\p{L}\\p{N}])(\\d{4})\\s+(году|г\\.)",
|
||||||
|
RegexOption.IGNORE_CASE
|
||||||
|
)
|
||||||
|
private val numberPattern = Regex(
|
||||||
|
"(?<![\\p{L}\\p{N}])([+-]?)(\\d{1,18})(?:[,.](\\d+))?(?![\\p{L}\\p{N}])"
|
||||||
|
)
|
||||||
|
|
||||||
|
fun normalize(text: String): String {
|
||||||
|
val prepared = text
|
||||||
|
.replace(Regex("(?<=\\d)[\\u00A0\\u202F](?=\\d{3}(?:\\D|$))"), "")
|
||||||
|
.replace(Regex("№\\s*(?=\\d)"), "номер ")
|
||||||
|
val yearsNormalized = yearInPrepositionalPattern.replace(prepared) { match ->
|
||||||
|
val year = match.groupValues[1].toInt()
|
||||||
|
"${yearInPrepositionalWords(year)} ${match.groupValues[2]}"
|
||||||
|
}
|
||||||
|
return numberPattern.replace(yearsNormalized) { match ->
|
||||||
|
val sign = when (match.groupValues[1]) {
|
||||||
|
"-" -> "минус "
|
||||||
|
"+" -> "плюс "
|
||||||
|
else -> ""
|
||||||
|
}
|
||||||
|
val integerDigits = match.groupValues[2]
|
||||||
|
val integer = if (integerDigits.length > 1 && integerDigits.startsWith('0')) {
|
||||||
|
integerDigits.toDigitWords()
|
||||||
|
} else {
|
||||||
|
integerDigits.toLongOrNull()?.let(::integerToWords) ?: integerDigits.toDigitWords()
|
||||||
|
}
|
||||||
|
val fractionDigits = match.groupValues[3]
|
||||||
|
if (fractionDigits.isBlank()) "$sign$integer" else "$sign$integer запятая ${fractionDigits.toDigitWords()}"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun yearInPrepositionalWords(value: Int): String {
|
||||||
|
require(value in 1000..9999)
|
||||||
|
if (value % 1000 == 0) return EXACT_THOUSANDTH_PREPOSITIONAL[value / 1000]
|
||||||
|
val thousands = value / 1000
|
||||||
|
val thousandsCardinal = underThousandToWords(thousands, feminine = true)
|
||||||
|
val thousandsWords = if (thousandsCardinal == "одна") {
|
||||||
|
SCALES[1]!!.formFor(thousands)
|
||||||
|
} else {
|
||||||
|
"$thousandsCardinal ${SCALES[1]!!.formFor(thousands)}"
|
||||||
|
}
|
||||||
|
return "$thousandsWords ${ordinalUnderThousandPrepositional(value % 1000)}"
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun ordinalUnderThousandPrepositional(value: Int): String {
|
||||||
|
require(value in 1..999)
|
||||||
|
val lastTwo = value % 100
|
||||||
|
val units = value % 10
|
||||||
|
return when {
|
||||||
|
units != 0 && lastTwo !in 11..19 ->
|
||||||
|
"${integerToWords((value - units).toLong())} ${ORDINAL_UNITS_PREPOSITIONAL[units]}".trim()
|
||||||
|
lastTwo in 11..19 ->
|
||||||
|
"${integerToWords((value - lastTwo).toLong())} ${ORDINAL_TEENS_PREPOSITIONAL[lastTwo - 10]}".trim()
|
||||||
|
lastTwo >= 20 ->
|
||||||
|
"${integerToWords((value - lastTwo).toLong())} ${ORDINAL_TENS_PREPOSITIONAL[lastTwo / 10]}".trim()
|
||||||
|
else -> ORDINAL_HUNDREDS_PREPOSITIONAL[value / 100]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun integerToWords(value: Long): String {
|
||||||
|
if (value == 0L) return DIGITS[0]
|
||||||
|
if (value < 0L) return "минус ${integerToWords(-value)}"
|
||||||
|
val groups = mutableListOf<Int>()
|
||||||
|
var remainder = value
|
||||||
|
while (remainder > 0) {
|
||||||
|
groups += (remainder % 1_000).toInt()
|
||||||
|
remainder /= 1_000
|
||||||
|
}
|
||||||
|
return groups.indices.reversed().mapNotNull { groupIndex ->
|
||||||
|
val group = groups[groupIndex]
|
||||||
|
if (group == 0) return@mapNotNull null
|
||||||
|
val feminine = groupIndex == 1
|
||||||
|
val words = underThousandToWords(group, feminine)
|
||||||
|
val scale = SCALES.getOrNull(groupIndex)
|
||||||
|
if (scale == null) words else "$words ${scale.formFor(group)}"
|
||||||
|
}.joinToString(" ")
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun underThousandToWords(value: Int, feminine: Boolean): String {
|
||||||
|
val words = mutableListOf<String>()
|
||||||
|
val hundreds = value / 100
|
||||||
|
if (hundreds > 0) words += HUNDREDS[hundreds]
|
||||||
|
val tail = value % 100
|
||||||
|
when {
|
||||||
|
tail in 10..19 -> words += TEENS[tail - 10]
|
||||||
|
else -> {
|
||||||
|
val tens = tail / 10
|
||||||
|
if (tens > 0) words += TENS[tens]
|
||||||
|
val units = tail % 10
|
||||||
|
if (units > 0) {
|
||||||
|
words += when {
|
||||||
|
feminine && units == 1 -> "одна"
|
||||||
|
feminine && units == 2 -> "две"
|
||||||
|
else -> DIGITS[units]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return words.joinToString(" ")
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun String.toDigitWords(): String = map { DIGITS[it.digitToInt()] }.joinToString(" ")
|
||||||
|
|
||||||
|
private data class Scale(val one: String, val few: String, val many: String) {
|
||||||
|
fun formFor(value: Int): String {
|
||||||
|
val lastTwo = value % 100
|
||||||
|
if (lastTwo in 11..14) return many
|
||||||
|
return when (value % 10) {
|
||||||
|
1 -> one
|
||||||
|
2, 3, 4 -> few
|
||||||
|
else -> many
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private val DIGITS = listOf("ноль", "один", "два", "три", "четыре", "пять", "шесть", "семь", "восемь", "девять")
|
||||||
|
private val TEENS = listOf(
|
||||||
|
"десять", "одиннадцать", "двенадцать", "тринадцать", "четырнадцать",
|
||||||
|
"пятнадцать", "шестнадцать", "семнадцать", "восемнадцать", "девятнадцать"
|
||||||
|
)
|
||||||
|
private val TENS = listOf("", "десять", "двадцать", "тридцать", "сорок", "пятьдесят", "шестьдесят", "семьдесят", "восемьдесят", "девяносто")
|
||||||
|
private val HUNDREDS = listOf("", "сто", "двести", "триста", "четыреста", "пятьсот", "шестьсот", "семьсот", "восемьсот", "девятьсот")
|
||||||
|
private val ORDINAL_UNITS_PREPOSITIONAL = listOf("", "первом", "втором", "третьем", "четвёртом", "пятом", "шестом", "седьмом", "восьмом", "девятом")
|
||||||
|
private val ORDINAL_TEENS_PREPOSITIONAL = listOf("десятом", "одиннадцатом", "двенадцатом", "тринадцатом", "четырнадцатом", "пятнадцатом", "шестнадцатом", "семнадцатом", "восемнадцатом", "девятнадцатом")
|
||||||
|
private val ORDINAL_TENS_PREPOSITIONAL = listOf("", "десятом", "двадцатом", "тридцатом", "сороковом", "пятидесятом", "шестидесятом", "семидесятом", "восьмидесятом", "девяностом")
|
||||||
|
private val ORDINAL_HUNDREDS_PREPOSITIONAL = listOf("", "сотом", "двухсотом", "трёхсотом", "четырёхсотом", "пятисотом", "шестисотом", "семисотом", "восьмисотом", "девятисотом")
|
||||||
|
private val EXACT_THOUSANDTH_PREPOSITIONAL = listOf("", "тысячном", "двухтысячном", "трёхтысячном", "четырёхтысячном", "пятитысячном", "шеститысячном", "семитысячном", "восьмитысячном", "девятитысячном")
|
||||||
|
private val SCALES = listOf(
|
||||||
|
null,
|
||||||
|
Scale("тысяча", "тысячи", "тысяч"),
|
||||||
|
Scale("миллион", "миллиона", "миллионов"),
|
||||||
|
Scale("миллиард", "миллиарда", "миллиардов"),
|
||||||
|
Scale("триллион", "триллиона", "триллионов"),
|
||||||
|
Scale("квадриллион", "квадриллиона", "квадриллионов")
|
||||||
|
)
|
||||||
|
}
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
package com.aletheia.app.ui.reader.tts
|
||||||
|
|
||||||
|
import java.util.Locale
|
||||||
|
|
||||||
|
internal object RussianSpeechTextSplitter {
|
||||||
|
fun splitSentences(text: String): List<String> {
|
||||||
|
val result = mutableListOf<String>()
|
||||||
|
var start = 0
|
||||||
|
var index = 0
|
||||||
|
while (index < text.length) {
|
||||||
|
val char = text[index]
|
||||||
|
val isEllipsis = char == '…' || text.startsWith("...", index)
|
||||||
|
val isTerminal = when {
|
||||||
|
char == '!' || char == '?' || isEllipsis -> true
|
||||||
|
char == '.' -> isSentencePeriod(text, index)
|
||||||
|
else -> false
|
||||||
|
}
|
||||||
|
if (!isTerminal) {
|
||||||
|
index += 1
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if (text.startsWith("...", index)) index += 2
|
||||||
|
var end = index + 1
|
||||||
|
while (end < text.length && text[end] in SENTENCE_CLOSERS) end += 1
|
||||||
|
val hasSentenceBoundary = end == text.length || text[end].isWhitespace()
|
||||||
|
if (hasSentenceBoundary && !isFollowedByDialogueDash(text, end)) {
|
||||||
|
val sentence = text.substring(start, end).trim()
|
||||||
|
if (sentence.isNotBlank()) result += sentence
|
||||||
|
start = end
|
||||||
|
while (start < text.length && text[start].isWhitespace()) start += 1
|
||||||
|
index = start
|
||||||
|
} else {
|
||||||
|
index = end
|
||||||
|
}
|
||||||
|
}
|
||||||
|
val remainder = text.substring(start).trim()
|
||||||
|
if (remainder.isNotBlank()) result += remainder
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun isFollowedByDialogueDash(text: String, end: Int): Boolean {
|
||||||
|
var index = end
|
||||||
|
while (index < text.length && text[index].isWhitespace()) index += 1
|
||||||
|
return text.getOrNull(index) in DIALOGUE_DASHES
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun isSentencePeriod(text: String, index: Int): Boolean {
|
||||||
|
if (text.getOrNull(index - 1)?.isDigit() == true && text.getOrNull(index + 1)?.isDigit() == true) return false
|
||||||
|
val prefix = text.substring(0, index)
|
||||||
|
val token = prefix.takeLastWhile { it.isLetter() }.lowercase(Locale.forLanguageTag("ru"))
|
||||||
|
if (token in NON_TERMINAL_ABBREVIATIONS) return false
|
||||||
|
if (token.length == 1 && prefix.lastOrNull()?.isUpperCase() == true) return false
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
private val SENTENCE_CLOSERS = setOf('"', '\'', '»', '”', ')')
|
||||||
|
private val DIALOGUE_DASHES = setOf('—', '–', '-')
|
||||||
|
private val NON_TERMINAL_ABBREVIATIONS = setOf(
|
||||||
|
"г", "гг", "ул", "стр", "рис", "им", "т", "д", "п", "н", "э", "е", "к",
|
||||||
|
"др", "см", "руб", "коп", "тыс", "млн", "млрд"
|
||||||
|
)
|
||||||
|
}
|
||||||
@@ -0,0 +1,708 @@
|
|||||||
|
package com.aletheia.app.ui.reader.tts
|
||||||
|
|
||||||
|
import ai.onnxruntime.OnnxTensor
|
||||||
|
import ai.onnxruntime.OrtEnvironment
|
||||||
|
import ai.onnxruntime.OrtSession
|
||||||
|
import android.content.Context
|
||||||
|
import android.media.AudioAttributes
|
||||||
|
import android.media.AudioFormat
|
||||||
|
import android.media.AudioTrack
|
||||||
|
import android.media.PlaybackParams
|
||||||
|
import android.util.Log
|
||||||
|
import java.io.ByteArrayOutputStream
|
||||||
|
import java.io.File
|
||||||
|
import java.io.FileOutputStream
|
||||||
|
import java.io.RandomAccessFile
|
||||||
|
import java.nio.FloatBuffer
|
||||||
|
import java.nio.LongBuffer
|
||||||
|
import java.nio.charset.StandardCharsets
|
||||||
|
import java.util.LinkedHashMap
|
||||||
|
import java.util.Locale
|
||||||
|
import java.util.concurrent.atomic.AtomicInteger
|
||||||
|
import kotlinx.coroutines.Dispatchers
|
||||||
|
import kotlinx.coroutines.ensureActive
|
||||||
|
import kotlinx.coroutines.withContext
|
||||||
|
import kotlin.coroutines.coroutineContext
|
||||||
|
|
||||||
|
class VoskSpeechEngine(context: Context) {
|
||||||
|
private val appContext = context.applicationContext
|
||||||
|
private val playbackGeneration = AtomicInteger()
|
||||||
|
private val trackLock = Any()
|
||||||
|
private val environment = OrtEnvironment.getEnvironment()
|
||||||
|
private var synthesisSession: OrtSession? = null
|
||||||
|
private var bertSession: OrtSession? = null
|
||||||
|
private var dictionary: PronunciationDictionary? = null
|
||||||
|
private var tokenizer: WordPieceTokenizer? = null
|
||||||
|
private var activeTrack: AudioTrack? = null
|
||||||
|
|
||||||
|
suspend fun initialize() = withContext(Dispatchers.IO) {
|
||||||
|
if (synthesisSession != null) return@withContext
|
||||||
|
val modelFile = copyAssetToInternalStorage(MODEL_ASSET, MODEL_FILE)
|
||||||
|
val bertFile = copyAssetToInternalStorage(BERT_ASSET, BERT_FILE)
|
||||||
|
val dictionaryFile = copyAssetToInternalStorage(DICTIONARY_ASSET, DICTIONARY_FILE)
|
||||||
|
val dictionaryIndex = appContext.assets.open(DICTIONARY_INDEX_ASSET).bufferedReader().useLines { lines ->
|
||||||
|
lines.mapNotNull { line ->
|
||||||
|
val separator = line.lastIndexOf('\t')
|
||||||
|
if (separator <= 0) null else DictionaryIndexEntry(
|
||||||
|
line.substring(0, separator),
|
||||||
|
line.substring(separator + 1).toLong()
|
||||||
|
)
|
||||||
|
}.toList()
|
||||||
|
}
|
||||||
|
val vocabulary = appContext.assets.open(VOCABULARY_ASSET).bufferedReader().useLines { lines ->
|
||||||
|
lines.mapIndexed { index, token -> token to index.toLong() }.toMap()
|
||||||
|
}
|
||||||
|
val options = OrtSession.SessionOptions().apply {
|
||||||
|
setOptimizationLevel(OrtSession.SessionOptions.OptLevel.NO_OPT)
|
||||||
|
setIntraOpNumThreads(INFERENCE_THREADS)
|
||||||
|
setInterOpNumThreads(1)
|
||||||
|
}
|
||||||
|
synthesisSession = environment.createSession(modelFile.absolutePath, options)
|
||||||
|
bertSession = environment.createSession(bertFile.absolutePath, options)
|
||||||
|
dictionary = PronunciationDictionary(dictionaryFile, dictionaryIndex)
|
||||||
|
tokenizer = WordPieceTokenizer(vocabulary)
|
||||||
|
Log.i(
|
||||||
|
TAG,
|
||||||
|
"Vosk TTS 0.9 loaded: model=${modelFile.length()} bert=${bertFile.length()} " +
|
||||||
|
"dictionary=${dictionaryFile.length()}"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
suspend fun speak(
|
||||||
|
text: String,
|
||||||
|
rate: Float,
|
||||||
|
pitch: Float,
|
||||||
|
pauseScale: Float,
|
||||||
|
articulationScale: Float
|
||||||
|
) = withContext(Dispatchers.Default) {
|
||||||
|
val localSynthesis = checkNotNull(synthesisSession) { "Vosk TTS is not initialized" }
|
||||||
|
val localBert = checkNotNull(bertSession) { "Vosk BERT is not initialized" }
|
||||||
|
val localDictionary = checkNotNull(dictionary) { "Vosk dictionary is not initialized" }
|
||||||
|
val localTokenizer = checkNotNull(tokenizer) { "Vosk tokenizer is not initialized" }
|
||||||
|
val generation = playbackGeneration.incrementAndGet()
|
||||||
|
for (chunk in splitText(text)) {
|
||||||
|
coroutineContext.ensureActive()
|
||||||
|
check(generation == playbackGeneration.get()) { "Playback stopped" }
|
||||||
|
val prepared = prepareInput(chunk.text, localDictionary, localTokenizer, localBert)
|
||||||
|
if (prepared.timeSteps <= 2) continue
|
||||||
|
Log.d(TAG, "Synthesizing ${chunk.text.length} chars as ${prepared.timeSteps} phonemes")
|
||||||
|
val output = synthesize(localSynthesis, prepared, articulationScale)
|
||||||
|
check(output.isNotEmpty()) { "Vosk returned empty audio" }
|
||||||
|
val playbackSpeed = (rate * BOOK_SPEED_SCALE).coerceIn(MIN_PLAYBACK_SPEED, MAX_PLAYBACK_SPEED)
|
||||||
|
val scaledPauseMs = if (chunk.trailingPauseMs == NO_TRAILING_PAUSE_MS) {
|
||||||
|
NO_TRAILING_PAUSE_MS
|
||||||
|
} else {
|
||||||
|
(chunk.trailingPauseMs * pauseScale / rate.coerceAtLeast(MIN_RATE_FOR_PAUSE_SCALING))
|
||||||
|
.toInt()
|
||||||
|
.coerceAtLeast(MIN_TRAILING_PAUSE_MS)
|
||||||
|
.coerceAtMost(chunk.maximumPauseMs)
|
||||||
|
}
|
||||||
|
Log.d(
|
||||||
|
TAG,
|
||||||
|
"Generated ${output.size} samples at $SAMPLE_RATE Hz for ${chunk.text.length} chars; " +
|
||||||
|
"speaker=$MALE_1_SPEAKER_ID speed=$playbackSpeed pitch=$pitch " +
|
||||||
|
"articulation=$articulationScale pause=${scaledPauseMs}ms"
|
||||||
|
)
|
||||||
|
play(output, generation, playbackSpeed, pitch, scaledPauseMs)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fun stop() {
|
||||||
|
playbackGeneration.incrementAndGet()
|
||||||
|
synchronized(trackLock) {
|
||||||
|
activeTrack?.runCatching { pause() }
|
||||||
|
activeTrack?.runCatching { flush() }
|
||||||
|
activeTrack?.runCatching { stop() }
|
||||||
|
activeTrack?.release()
|
||||||
|
activeTrack = null
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fun destroy() {
|
||||||
|
stop()
|
||||||
|
dictionary?.close()
|
||||||
|
dictionary = null
|
||||||
|
tokenizer = null
|
||||||
|
bertSession?.close()
|
||||||
|
bertSession = null
|
||||||
|
synthesisSession?.close()
|
||||||
|
synthesisSession = null
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun prepareInput(
|
||||||
|
text: String,
|
||||||
|
dictionary: PronunciationDictionary,
|
||||||
|
tokenizer: WordPieceTokenizer,
|
||||||
|
bert: OrtSession
|
||||||
|
): PreparedInput {
|
||||||
|
val normalized = RussianNumberNormalizer.normalize(text)
|
||||||
|
.replace(Regex("([А-Яа-яЁё])\\u0301"), "+$1")
|
||||||
|
.lowercase(Locale.forLanguageTag("ru"))
|
||||||
|
.replace('—', '-')
|
||||||
|
.replace('–', '-')
|
||||||
|
.replace("…", "...")
|
||||||
|
.replace(Regex("[«»“”]"), "\"")
|
||||||
|
.replace(Regex("\\s+"), " ")
|
||||||
|
.trim()
|
||||||
|
val encoding = tokenizer.encode(normalized)
|
||||||
|
val tokenEmbeddings = runBert(bert, encoding.ids)
|
||||||
|
val selectedEmbeddings = encoding.embeddingPositions.map { position -> tokenEmbeddings[position] }
|
||||||
|
val phones = buildMultistreamPhones(normalized, dictionary)
|
||||||
|
val timeSteps = phones.size
|
||||||
|
val features = LongArray(FEATURE_CHANNELS * timeSteps)
|
||||||
|
val embeddings = FloatArray(BERT_DIMENSIONS * timeSteps)
|
||||||
|
phones.forEachIndexed { time, phone ->
|
||||||
|
phone.features.forEachIndexed { channel, value -> features[channel * timeSteps + time] = value }
|
||||||
|
val embedding = selectedEmbeddings[phone.bertWordIndex.coerceIn(selectedEmbeddings.indices)]
|
||||||
|
for (channel in 0 until BERT_DIMENSIONS) {
|
||||||
|
embeddings[channel * timeSteps + time] = embedding[channel]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return PreparedInput(features, embeddings, timeSteps)
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun runBert(session: OrtSession, ids: LongArray): Array<FloatArray> {
|
||||||
|
val shape = longArrayOf(1, ids.size.toLong())
|
||||||
|
val inputIds = OnnxTensor.createTensor(environment, LongBuffer.wrap(ids), shape)
|
||||||
|
val attentionMask = OnnxTensor.createTensor(
|
||||||
|
environment,
|
||||||
|
LongBuffer.wrap(LongArray(ids.size) { 1L }),
|
||||||
|
shape
|
||||||
|
)
|
||||||
|
val tokenTypes = OnnxTensor.createTensor(environment, LongBuffer.wrap(LongArray(ids.size)), shape)
|
||||||
|
try {
|
||||||
|
session.run(
|
||||||
|
mapOf(
|
||||||
|
"input_ids" to inputIds,
|
||||||
|
"attention_mask" to attentionMask,
|
||||||
|
"token_type_ids" to tokenTypes
|
||||||
|
)
|
||||||
|
).use { result ->
|
||||||
|
@Suppress("UNCHECKED_CAST")
|
||||||
|
return result[0].value as Array<FloatArray>
|
||||||
|
}
|
||||||
|
} finally {
|
||||||
|
inputIds.close()
|
||||||
|
attentionMask.close()
|
||||||
|
tokenTypes.close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun synthesize(session: OrtSession, prepared: PreparedInput, articulationScale: Float): FloatArray {
|
||||||
|
val input = OnnxTensor.createTensor(
|
||||||
|
environment,
|
||||||
|
LongBuffer.wrap(prepared.features),
|
||||||
|
longArrayOf(1, FEATURE_CHANNELS.toLong(), prepared.timeSteps.toLong())
|
||||||
|
)
|
||||||
|
val inputLengths = OnnxTensor.createTensor(
|
||||||
|
environment,
|
||||||
|
LongBuffer.wrap(longArrayOf(prepared.timeSteps.toLong())),
|
||||||
|
longArrayOf(1)
|
||||||
|
)
|
||||||
|
val scales = OnnxTensor.createTensor(
|
||||||
|
environment,
|
||||||
|
FloatBuffer.wrap(
|
||||||
|
floatArrayOf(
|
||||||
|
NOISE_SCALE,
|
||||||
|
articulationScale.coerceIn(MIN_ARTICULATION_SCALE, MAX_ARTICULATION_SCALE),
|
||||||
|
DURATION_NOISE_SCALE
|
||||||
|
)
|
||||||
|
),
|
||||||
|
longArrayOf(3)
|
||||||
|
)
|
||||||
|
val speaker = OnnxTensor.createTensor(
|
||||||
|
environment,
|
||||||
|
LongBuffer.wrap(longArrayOf(MALE_1_SPEAKER_ID.toLong())),
|
||||||
|
longArrayOf(1)
|
||||||
|
)
|
||||||
|
val bert = OnnxTensor.createTensor(
|
||||||
|
environment,
|
||||||
|
FloatBuffer.wrap(prepared.bertEmbeddings),
|
||||||
|
longArrayOf(1, BERT_DIMENSIONS.toLong(), prepared.timeSteps.toLong())
|
||||||
|
)
|
||||||
|
try {
|
||||||
|
session.run(
|
||||||
|
mapOf(
|
||||||
|
"input" to input,
|
||||||
|
"input_lengths" to inputLengths,
|
||||||
|
"scales" to scales,
|
||||||
|
"sid" to speaker,
|
||||||
|
"bert" to bert
|
||||||
|
)
|
||||||
|
).use { result ->
|
||||||
|
return flattenFloatOutput(result[0].value)
|
||||||
|
}
|
||||||
|
} finally {
|
||||||
|
input.close()
|
||||||
|
inputLengths.close()
|
||||||
|
scales.close()
|
||||||
|
speaker.close()
|
||||||
|
bert.close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun buildMultistreamPhones(
|
||||||
|
text: String,
|
||||||
|
dictionary: PronunciationDictionary
|
||||||
|
): List<PhoneFeatures> {
|
||||||
|
val rawPhones = mutableListOf(RawPhone("^", emptyList(), 0, 0))
|
||||||
|
val word = StringBuilder()
|
||||||
|
val pendingPunctuation = mutableListOf<String>()
|
||||||
|
var inQuote = 0
|
||||||
|
var bertWordIndex = 1
|
||||||
|
fun flushWord() {
|
||||||
|
if (word.isEmpty()) return
|
||||||
|
val value = word.toString()
|
||||||
|
val pronunciation = dictionary.find(value) ?: RussianG2p.convert(value)
|
||||||
|
pronunciation.split(' ').filter(String::isNotBlank).forEach { phone ->
|
||||||
|
rawPhones += RawPhone(phone, emptyList(), inQuote, bertWordIndex)
|
||||||
|
}
|
||||||
|
word.clear()
|
||||||
|
bertWordIndex += 1
|
||||||
|
}
|
||||||
|
fun appendSpace() {
|
||||||
|
rawPhones += RawPhone(" ", pendingPunctuation.toList(), inQuote, bertWordIndex)
|
||||||
|
pendingPunctuation.clear()
|
||||||
|
}
|
||||||
|
var index = 0
|
||||||
|
while (index < text.length) {
|
||||||
|
if (text.startsWith("...", index)) {
|
||||||
|
flushWord()
|
||||||
|
pendingPunctuation += "..."
|
||||||
|
index += 3
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
val char = text[index]
|
||||||
|
when {
|
||||||
|
char == '"' || char == '«' || char == '»' || char == '“' || char == '”' -> {
|
||||||
|
flushWord()
|
||||||
|
inQuote = 1 - inQuote
|
||||||
|
}
|
||||||
|
char.isWhitespace() -> {
|
||||||
|
flushWord()
|
||||||
|
appendSpace()
|
||||||
|
}
|
||||||
|
char == '-' && text.getOrNull(index - 1)?.isLetterOrDigit() == true &&
|
||||||
|
text.getOrNull(index + 1)?.isLetterOrDigit() == true -> {
|
||||||
|
// A lexical hyphen separates WordPiece words but must not create an acoustic pause.
|
||||||
|
flushWord()
|
||||||
|
}
|
||||||
|
char in MULTISTREAM_PUNCTUATION -> {
|
||||||
|
flushWord()
|
||||||
|
pendingPunctuation += char.toString()
|
||||||
|
}
|
||||||
|
char.isLetterOrDigit() || char == '+' -> word.append(char)
|
||||||
|
}
|
||||||
|
index += 1
|
||||||
|
}
|
||||||
|
flushWord()
|
||||||
|
appendSpace()
|
||||||
|
rawPhones += RawPhone("$", emptyList(), 0, bertWordIndex)
|
||||||
|
|
||||||
|
var lastPunctuation = " "
|
||||||
|
var lastSentencePunctuation = " "
|
||||||
|
val reversed = mutableListOf<PhoneFeatures>()
|
||||||
|
for (phone in rawPhones.asReversed()) {
|
||||||
|
lastSentencePunctuation = when {
|
||||||
|
"..." in phone.punctuation -> "..."
|
||||||
|
"." in phone.punctuation -> "."
|
||||||
|
"!" in phone.punctuation -> "!"
|
||||||
|
"?" in phone.punctuation -> "?"
|
||||||
|
"-" in phone.punctuation -> "-"
|
||||||
|
else -> lastSentencePunctuation
|
||||||
|
}
|
||||||
|
if (phone.punctuation.isNotEmpty()) lastPunctuation = phone.punctuation.first()
|
||||||
|
val currentPunctuation = phone.punctuation.firstOrNull() ?: "_"
|
||||||
|
val values = longArrayOf(
|
||||||
|
PHONEME_IDS.getValue(phone.phone),
|
||||||
|
PHONEME_IDS.getValue(currentPunctuation),
|
||||||
|
phone.inQuote.toLong(),
|
||||||
|
PHONEME_IDS.getValue(lastPunctuation),
|
||||||
|
PHONEME_IDS.getValue(lastSentencePunctuation)
|
||||||
|
)
|
||||||
|
reversed += PhoneFeatures(values, phone.bertWordIndex)
|
||||||
|
}
|
||||||
|
return reversed.asReversed()
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun flattenFloatOutput(value: Any?): FloatArray = when (value) {
|
||||||
|
is FloatArray -> value
|
||||||
|
is Array<*> -> value.asSequence().flatMap { flattenFloatOutput(it).asSequence() }.toList().toFloatArray()
|
||||||
|
else -> error("Unexpected Vosk output type: ${value?.javaClass?.name}")
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun play(
|
||||||
|
samples: FloatArray,
|
||||||
|
generation: Int,
|
||||||
|
playbackSpeed: Float,
|
||||||
|
pitch: Float,
|
||||||
|
trailingPauseMs: Int
|
||||||
|
) {
|
||||||
|
val minBytes = AudioTrack.getMinBufferSize(
|
||||||
|
SAMPLE_RATE,
|
||||||
|
AudioFormat.CHANNEL_OUT_MONO,
|
||||||
|
AudioFormat.ENCODING_PCM_FLOAT
|
||||||
|
).coerceAtLeast(4096)
|
||||||
|
val track = AudioTrack.Builder()
|
||||||
|
.setAudioAttributes(
|
||||||
|
AudioAttributes.Builder()
|
||||||
|
.setUsage(AudioAttributes.USAGE_MEDIA)
|
||||||
|
.setContentType(AudioAttributes.CONTENT_TYPE_SPEECH)
|
||||||
|
.build()
|
||||||
|
)
|
||||||
|
.setAudioFormat(
|
||||||
|
AudioFormat.Builder()
|
||||||
|
.setEncoding(AudioFormat.ENCODING_PCM_FLOAT)
|
||||||
|
.setSampleRate(SAMPLE_RATE)
|
||||||
|
.setChannelMask(AudioFormat.CHANNEL_OUT_MONO)
|
||||||
|
.build()
|
||||||
|
)
|
||||||
|
.setBufferSizeInBytes(minBytes)
|
||||||
|
.setTransferMode(AudioTrack.MODE_STREAM)
|
||||||
|
.build()
|
||||||
|
track.playbackParams = PlaybackParams()
|
||||||
|
.allowDefaults()
|
||||||
|
.setSpeed(playbackSpeed)
|
||||||
|
.setPitch(pitch.coerceIn(MIN_PLAYBACK_PITCH, MAX_PLAYBACK_PITCH))
|
||||||
|
synchronized(trackLock) { activeTrack = track }
|
||||||
|
try {
|
||||||
|
track.play()
|
||||||
|
var offset = 0
|
||||||
|
while (offset < samples.size && generation == playbackGeneration.get()) {
|
||||||
|
val written = track.write(samples, offset, samples.size - offset, AudioTrack.WRITE_BLOCKING)
|
||||||
|
if (written <= 0) error("AudioTrack write failed: $written")
|
||||||
|
offset += written
|
||||||
|
}
|
||||||
|
var silenceSamples = trailingPauseMs * SAMPLE_RATE / 1_000
|
||||||
|
while (silenceSamples > 0 && generation == playbackGeneration.get()) {
|
||||||
|
val count = minOf(silenceSamples, SILENCE_BUFFER.size)
|
||||||
|
val written = track.write(SILENCE_BUFFER, 0, count, AudioTrack.WRITE_BLOCKING)
|
||||||
|
if (written <= 0) error("AudioTrack silence write failed: $written")
|
||||||
|
silenceSamples -= written
|
||||||
|
}
|
||||||
|
check(generation == playbackGeneration.get()) { "Playback stopped" }
|
||||||
|
} finally {
|
||||||
|
synchronized(trackLock) { if (activeTrack === track) activeTrack = null }
|
||||||
|
track.runCatching { stop() }
|
||||||
|
track.release()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun splitText(text: String): List<SpeechChunk> {
|
||||||
|
val normalized = text
|
||||||
|
.replace(Regex("\\[\\d+(?:[,–-]\\d+)*\\]"), " ")
|
||||||
|
.replace("\r\n", "\n")
|
||||||
|
.replace('\r', '\n')
|
||||||
|
.replace(Regex("[\\t\\x0B\\f ]+"), " ")
|
||||||
|
.replace(Regex(" *\\n+ *"), "\n")
|
||||||
|
.trim()
|
||||||
|
if (normalized.isBlank()) return emptyList()
|
||||||
|
val chunks = mutableListOf<SpeechChunk>()
|
||||||
|
for (paragraph in normalized.split('\n').filter(String::isNotBlank)) {
|
||||||
|
for (sentence in RussianSpeechTextSplitter.splitSentences(paragraph)) {
|
||||||
|
for (unit in splitControlledPauses(sentence)) {
|
||||||
|
val parts = splitAtWordBoundaries(unit.text)
|
||||||
|
parts.forEachIndexed { index, part ->
|
||||||
|
val isLastPart = index == parts.lastIndex
|
||||||
|
chunks += SpeechChunk(
|
||||||
|
text = part.trim(),
|
||||||
|
trailingPauseMs = if (isLastPart) unit.trailingPauseMs else NO_TRAILING_PAUSE_MS,
|
||||||
|
maximumPauseMs = if (isLastPart) unit.maximumPauseMs else NO_TRAILING_PAUSE_MS
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return chunks.filter { it.text.isNotBlank() }
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun splitControlledPauses(text: String): List<ProsodicUnit> {
|
||||||
|
val units = mutableListOf<ProsodicUnit>()
|
||||||
|
val current = StringBuilder()
|
||||||
|
fun flushControlledPause() {
|
||||||
|
val value = current.toString().trim()
|
||||||
|
if (value.isNotBlank()) {
|
||||||
|
units += ProsodicUnit(value, CONTROLLED_PUNCTUATION_PAUSE_MS, CONTROLLED_PUNCTUATION_MAX_PAUSE_MS)
|
||||||
|
}
|
||||||
|
current.clear()
|
||||||
|
}
|
||||||
|
text.forEachIndexed { index, char ->
|
||||||
|
val isWordInternalHyphen = char == '-' &&
|
||||||
|
text.getOrNull(index - 1)?.isLetterOrDigit() == true &&
|
||||||
|
text.getOrNull(index + 1)?.isLetterOrDigit() == true
|
||||||
|
when {
|
||||||
|
char in CONTROLLED_PAUSE_PUNCTUATION && !isWordInternalHyphen -> flushControlledPause()
|
||||||
|
else -> current.append(char)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
val remainder = current.toString().trim()
|
||||||
|
if (remainder.isNotBlank()) units += ProsodicUnit(remainder, pauseAfter(remainder), Int.MAX_VALUE)
|
||||||
|
return units
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun splitAtWordBoundaries(text: String): List<String> {
|
||||||
|
if (text.length <= MAX_CHUNK_CHARS) return listOf(text)
|
||||||
|
val parts = mutableListOf<String>()
|
||||||
|
var remaining = text.trim()
|
||||||
|
while (remaining.length > MAX_CHUNK_CHARS) {
|
||||||
|
val boundary = remaining.lastIndexOf(' ', startIndex = MAX_CHUNK_CHARS)
|
||||||
|
.takeIf { it > MIN_CHUNK_CHARS }
|
||||||
|
?: remaining.indexOf(' ', startIndex = MAX_CHUNK_CHARS).takeIf { it >= 0 }
|
||||||
|
?: remaining.length
|
||||||
|
parts += remaining.substring(0, boundary).trim()
|
||||||
|
remaining = remaining.substring(boundary).trimStart()
|
||||||
|
}
|
||||||
|
if (remaining.isNotBlank()) parts += remaining
|
||||||
|
return parts
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun pauseAfter(text: String): Int = when (text.trimEnd().lastOrNull()) {
|
||||||
|
'…' -> ELLIPSIS_PAUSE_MS
|
||||||
|
'?' -> QUESTION_PAUSE_MS
|
||||||
|
'!' -> EXCLAMATION_PAUSE_MS
|
||||||
|
'.' -> SENTENCE_PAUSE_MS
|
||||||
|
':', ';' -> CLAUSE_PAUSE_MS
|
||||||
|
else -> NO_TRAILING_PAUSE_MS
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun copyAssetToInternalStorage(assetPath: String, fileName: String): File {
|
||||||
|
val destination = File(appContext.filesDir, "tts/vosk_0_9/$fileName")
|
||||||
|
val expectedSize = appContext.assets.open(assetPath).use { input ->
|
||||||
|
var total = 0L
|
||||||
|
val buffer = ByteArray(COPY_BUFFER_SIZE)
|
||||||
|
while (true) {
|
||||||
|
val read = input.read(buffer)
|
||||||
|
if (read < 0) break
|
||||||
|
total += read
|
||||||
|
}
|
||||||
|
total
|
||||||
|
}
|
||||||
|
if (destination.isFile && destination.length() == expectedSize) return destination
|
||||||
|
destination.parentFile?.mkdirs()
|
||||||
|
val temporary = File(destination.parentFile, "$fileName.tmp")
|
||||||
|
appContext.assets.open(assetPath).use { input ->
|
||||||
|
FileOutputStream(temporary).use { output -> input.copyTo(output, COPY_BUFFER_SIZE) }
|
||||||
|
}
|
||||||
|
check(temporary.length() == expectedSize) { "Vosk asset copy is incomplete: $fileName" }
|
||||||
|
if (destination.exists()) check(destination.delete()) { "Cannot replace old Vosk asset: $fileName" }
|
||||||
|
check(temporary.renameTo(destination)) { "Cannot install Vosk asset: $fileName" }
|
||||||
|
return destination
|
||||||
|
}
|
||||||
|
|
||||||
|
private class WordPieceTokenizer(private val vocabulary: Map<String, Long>) {
|
||||||
|
fun encode(text: String): BertEncoding {
|
||||||
|
val ids = mutableListOf(vocabulary.getValue("[CLS]"))
|
||||||
|
val embeddingPositions = mutableListOf(0)
|
||||||
|
for (token in basicTokens(text)) {
|
||||||
|
val punctuation = token in BERT_EXCLUDED_PUNCTUATION
|
||||||
|
val pieces = if (punctuation) {
|
||||||
|
listOf(vocabulary[token] ?: vocabulary.getValue("[UNK]"))
|
||||||
|
} else {
|
||||||
|
wordPieces(token)
|
||||||
|
}
|
||||||
|
if (!punctuation) embeddingPositions += ids.size
|
||||||
|
ids += pieces
|
||||||
|
}
|
||||||
|
ids += vocabulary.getValue("[SEP]")
|
||||||
|
embeddingPositions += ids.lastIndex
|
||||||
|
return BertEncoding(ids.toLongArray(), embeddingPositions.toIntArray())
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun basicTokens(text: String): List<String> {
|
||||||
|
val result = mutableListOf<String>()
|
||||||
|
val word = StringBuilder()
|
||||||
|
fun flush() {
|
||||||
|
if (word.isNotEmpty()) result += word.toString().also { word.clear() }
|
||||||
|
}
|
||||||
|
text.forEach { char ->
|
||||||
|
when {
|
||||||
|
char.isWhitespace() -> flush()
|
||||||
|
char.isLetterOrDigit() -> word.append(char.lowercaseChar())
|
||||||
|
char == '+' -> Unit // Explicit stress marker belongs to G2P, not to BERT tokenization.
|
||||||
|
else -> {
|
||||||
|
flush()
|
||||||
|
result += char.toString()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
flush()
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun wordPieces(token: String): List<Long> {
|
||||||
|
if (token.length > MAX_WORD_PIECE_CHARS) return listOf(vocabulary.getValue("[UNK]"))
|
||||||
|
val result = mutableListOf<Long>()
|
||||||
|
var start = 0
|
||||||
|
while (start < token.length) {
|
||||||
|
var end = token.length
|
||||||
|
var found: Long? = null
|
||||||
|
while (start < end) {
|
||||||
|
val piece = token.substring(start, end).let { if (start == 0) it else "##$it" }
|
||||||
|
found = vocabulary[piece]
|
||||||
|
if (found != null) break
|
||||||
|
end -= 1
|
||||||
|
}
|
||||||
|
if (found == null) return listOf(vocabulary.getValue("[UNK]"))
|
||||||
|
result += found
|
||||||
|
start = end
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private data class DictionaryIndexEntry(val word: String, val offset: Long)
|
||||||
|
|
||||||
|
private class PronunciationDictionary(
|
||||||
|
file: File,
|
||||||
|
private val index: List<DictionaryIndexEntry>
|
||||||
|
) : AutoCloseable {
|
||||||
|
private val source = RandomAccessFile(file, "r")
|
||||||
|
private val cache = object : LinkedHashMap<String, String>(CACHE_SIZE, 0.75f, true) {
|
||||||
|
override fun removeEldestEntry(eldest: MutableMap.MutableEntry<String, String>?): Boolean = size > CACHE_SIZE
|
||||||
|
}
|
||||||
|
|
||||||
|
@Synchronized
|
||||||
|
fun find(word: String): String? {
|
||||||
|
cache[word]?.let { return it }
|
||||||
|
var low = 0
|
||||||
|
var high = index.lastIndex
|
||||||
|
while (low <= high) {
|
||||||
|
val midpoint = (low + high) ushr 1
|
||||||
|
if (index[midpoint].word <= word) low = midpoint + 1 else high = midpoint - 1
|
||||||
|
}
|
||||||
|
val rangeIndex = high.coerceAtLeast(0)
|
||||||
|
val rangeEnd = index.getOrNull(rangeIndex + 1)?.offset ?: source.length()
|
||||||
|
source.seek(index[rangeIndex].offset)
|
||||||
|
while (source.filePointer < rangeEnd) {
|
||||||
|
val line = readUtf8Line(source) ?: break
|
||||||
|
val entryWord = line.substringBefore('\t')
|
||||||
|
when {
|
||||||
|
entryWord < word -> continue
|
||||||
|
entryWord > word -> return null
|
||||||
|
else -> {
|
||||||
|
val pronunciation = line.substringAfter('\t', "").trim()
|
||||||
|
if (pronunciation.isNotEmpty()) cache[word] = pronunciation
|
||||||
|
return pronunciation.takeIf(String::isNotEmpty)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return null
|
||||||
|
}
|
||||||
|
|
||||||
|
override fun close() = source.close()
|
||||||
|
|
||||||
|
private fun readUtf8Line(file: RandomAccessFile): String? {
|
||||||
|
val bytes = ByteArrayOutputStream(64)
|
||||||
|
while (true) {
|
||||||
|
val value = file.read()
|
||||||
|
if (value < 0) return if (bytes.size() == 0) null else bytes.toString(StandardCharsets.UTF_8.name())
|
||||||
|
if (value == '\n'.code) break
|
||||||
|
if (value != '\r'.code) bytes.write(value)
|
||||||
|
}
|
||||||
|
return bytes.toString(StandardCharsets.UTF_8.name())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private object RussianG2p {
|
||||||
|
private val softLetters = setOf('я', 'ё', 'ю', 'и', 'ь', 'е')
|
||||||
|
private val syllableStarts = setOf('#', 'ъ', 'ь', 'а', 'я', 'о', 'ё', 'у', 'ю', 'э', 'е', 'и', 'ы', '-')
|
||||||
|
private val softHardConsonants = mapOf(
|
||||||
|
'б' to "b", 'в' to "v", 'г' to "g", 'д' to "d", 'з' to "z", 'к' to "k",
|
||||||
|
'л' to "l", 'м' to "m", 'н' to "n", 'п' to "p", 'р' to "r", 'с' to "s",
|
||||||
|
'т' to "t", 'ф' to "f", 'х' to "h"
|
||||||
|
)
|
||||||
|
private val otherConsonants = mapOf('ж' to "zh", 'ц' to "c", 'ч' to "ch", 'ш' to "sh", 'щ' to "sch", 'й' to "j")
|
||||||
|
private val vowels = mapOf('а' to "a", 'я' to "a", 'у' to "u", 'ю' to "u", 'о' to "o", 'ё' to "o", 'э' to "e", 'е' to "e", 'и' to "i", 'ы' to "y")
|
||||||
|
|
||||||
|
fun convert(word: String): String {
|
||||||
|
val marked = mutableListOf<Pair<Char, Int>>()
|
||||||
|
var stress = 0
|
||||||
|
for (char in "#$word#") {
|
||||||
|
if (char == '+') stress = 1 else {
|
||||||
|
marked += char to stress
|
||||||
|
stress = 0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
val phones = mutableListOf<String>()
|
||||||
|
for (index in marked.indices) {
|
||||||
|
val (char, accent) = marked[index]
|
||||||
|
val next = marked.getOrNull(index + 1)?.first
|
||||||
|
when {
|
||||||
|
char in softHardConsonants -> phones += softHardConsonants.getValue(char) + if (next in softLetters) "j" else ""
|
||||||
|
char in otherConsonants -> phones += otherConsonants.getValue(char)
|
||||||
|
char in vowels -> {
|
||||||
|
val previous = marked.getOrNull(index - 1)?.first ?: '#'
|
||||||
|
if (previous in syllableStarts && char in setOf('я', 'ю', 'е', 'ё')) phones += "j"
|
||||||
|
phones += vowels.getValue(char) + accent
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return phones.joinToString(" ")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private companion object {
|
||||||
|
data class PreparedInput(
|
||||||
|
val features: LongArray,
|
||||||
|
val bertEmbeddings: FloatArray,
|
||||||
|
val timeSteps: Int
|
||||||
|
)
|
||||||
|
|
||||||
|
data class BertEncoding(val ids: LongArray, val embeddingPositions: IntArray)
|
||||||
|
data class RawPhone(val phone: String, val punctuation: List<String>, val inQuote: Int, val bertWordIndex: Int)
|
||||||
|
data class PhoneFeatures(val features: LongArray, val bertWordIndex: Int)
|
||||||
|
data class SpeechChunk(val text: String, val trailingPauseMs: Int, val maximumPauseMs: Int = Int.MAX_VALUE)
|
||||||
|
data class ProsodicUnit(val text: String, val trailingPauseMs: Int, val maximumPauseMs: Int)
|
||||||
|
|
||||||
|
const val TAG = "AletheiaVoskTts"
|
||||||
|
const val MODEL_FILE = "model.onnx"
|
||||||
|
const val BERT_FILE = "bert.int8.onnx"
|
||||||
|
const val DICTIONARY_FILE = "dict.tsv"
|
||||||
|
const val MODEL_ASSET = "tts/vosk_0_9/$MODEL_FILE"
|
||||||
|
const val BERT_ASSET = "tts/vosk_0_9/$BERT_FILE"
|
||||||
|
const val DICTIONARY_ASSET = "tts/vosk_0_9/$DICTIONARY_FILE"
|
||||||
|
const val DICTIONARY_INDEX_ASSET = "tts/vosk_0_9/dictionary.index"
|
||||||
|
const val VOCABULARY_ASSET = "tts/vosk_0_9/vocab.txt"
|
||||||
|
const val SAMPLE_RATE = 22_050
|
||||||
|
const val MALE_1_SPEAKER_ID = 4
|
||||||
|
const val FEATURE_CHANNELS = 5
|
||||||
|
const val BERT_DIMENSIONS = 768
|
||||||
|
const val INFERENCE_THREADS = 4
|
||||||
|
const val COPY_BUFFER_SIZE = 1024 * 1024
|
||||||
|
const val CACHE_SIZE = 4_096
|
||||||
|
const val MAX_WORD_PIECE_CHARS = 100
|
||||||
|
const val MAX_CHUNK_CHARS = 120
|
||||||
|
const val MIN_CHUNK_CHARS = 80
|
||||||
|
const val BOOK_SPEED_SCALE = 0.82f
|
||||||
|
const val MIN_RATE_FOR_PAUSE_SCALING = 0.10f
|
||||||
|
const val MIN_PLAYBACK_SPEED = 0.50f
|
||||||
|
const val MAX_PLAYBACK_SPEED = 1.50f
|
||||||
|
const val MIN_PLAYBACK_PITCH = 0.75f
|
||||||
|
const val MAX_PLAYBACK_PITCH = 1.30f
|
||||||
|
const val MIN_ARTICULATION_SCALE = 0.90f
|
||||||
|
const val MAX_ARTICULATION_SCALE = 1.20f
|
||||||
|
const val NOISE_SCALE = 0.8f
|
||||||
|
const val DURATION_NOISE_SCALE = 0.8f
|
||||||
|
const val NO_TRAILING_PAUSE_MS = 0
|
||||||
|
const val MIN_TRAILING_PAUSE_MS = 80
|
||||||
|
const val CLAUSE_PAUSE_MS = 320
|
||||||
|
const val CONTROLLED_PUNCTUATION_PAUSE_MS = 300
|
||||||
|
const val CONTROLLED_PUNCTUATION_MAX_PAUSE_MS = 300
|
||||||
|
const val SENTENCE_PAUSE_MS = 600
|
||||||
|
const val EXCLAMATION_PAUSE_MS = 600
|
||||||
|
const val QUESTION_PAUSE_MS = 600
|
||||||
|
const val ELLIPSIS_PAUSE_MS = 1_200
|
||||||
|
|
||||||
|
val SILENCE_BUFFER = FloatArray(SAMPLE_RATE / 10)
|
||||||
|
val CONTROLLED_PAUSE_PUNCTUATION = setOf('—', '–', '-', '(', ')', ';', ':')
|
||||||
|
val MULTISTREAM_PUNCTUATION = setOf('!', '(', ')', ',', '-', '.', ':', ';', '?')
|
||||||
|
val BERT_EXCLUDED_PUNCTUATION = setOf("-", ",", ".", "?", "!", ";", ":", "\"")
|
||||||
|
val PHONEME_IDS = listOf(
|
||||||
|
"_", "^", "$", " ", "!", "'", "(", ")", ",", "-", ".", "...", ":", ";", "?",
|
||||||
|
"a0", "a1", "b", "bj", "c", "ch", "d", "dj", "e0", "e1", "f", "fj", "g", "gj",
|
||||||
|
"h", "hj", "i0", "i1", "j", "k", "kj", "l", "lj", "m", "mj", "n", "nj", "o0",
|
||||||
|
"o1", "p", "pj", "r", "rj", "s", "sch", "sh", "sj", "t", "tj", "u0", "u1", "v",
|
||||||
|
"vj", "y0", "y1", "z", "zh", "zj"
|
||||||
|
).mapIndexed { index, phoneme -> phoneme to index.toLong() }.toMap()
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,40 @@
|
|||||||
|
package com.aletheia.app.ui.reader.tts
|
||||||
|
|
||||||
|
import org.junit.Assert.assertEquals
|
||||||
|
import org.junit.Test
|
||||||
|
|
||||||
|
class RussianNumberNormalizerTest {
|
||||||
|
@Test
|
||||||
|
fun `normalizes cardinal integers`() {
|
||||||
|
assertEquals("ноль один двадцать один", RussianNumberNormalizer.normalize("0 1 21"))
|
||||||
|
assertEquals("одна тысяча один", RussianNumberNormalizer.normalize("1001"))
|
||||||
|
assertEquals("две тысячи двадцать шесть", RussianNumberNormalizer.normalize("2026"))
|
||||||
|
assertEquals("двадцать две тысячи", RussianNumberNormalizer.normalize("22000"))
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test
|
||||||
|
fun `preserves leading zeroes and reads decimal digits`() {
|
||||||
|
assertEquals("ноль ноль семь", RussianNumberNormalizer.normalize("007"))
|
||||||
|
assertEquals("минус двенадцать запятая ноль пять", RussianNumberNormalizer.normalize("-12,05"))
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test
|
||||||
|
fun `normalizes number sign and non breaking groups`() {
|
||||||
|
assertEquals("номер пять", RussianNumberNormalizer.normalize("№ 5"))
|
||||||
|
assertEquals("двенадцать тысяч триста сорок пять", RussianNumberNormalizer.normalize("12\u00A0345"))
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test
|
||||||
|
fun `does not rewrite digits inside words`() {
|
||||||
|
assertEquals("глава2 и abc123", RussianNumberNormalizer.normalize("глава2 и abc123"))
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test
|
||||||
|
fun `inflects four digit years before году`() {
|
||||||
|
assertEquals(
|
||||||
|
"В две тысячи двадцать шестом году и в тысяча восемьсот девяносто седьмом году",
|
||||||
|
RussianNumberNormalizer.normalize("В 2026 году и в 1897 году")
|
||||||
|
)
|
||||||
|
assertEquals("В двухтысячном году", RussianNumberNormalizer.normalize("В 2000 году"))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,27 @@
|
|||||||
|
package com.aletheia.app.ui.reader.tts
|
||||||
|
|
||||||
|
import org.junit.Assert.assertEquals
|
||||||
|
import org.junit.Test
|
||||||
|
|
||||||
|
class RussianSpeechTextSplitterTest {
|
||||||
|
@Test
|
||||||
|
fun keepsTerminalPunctuationAndDialogueDashInOneSpeechSentence() {
|
||||||
|
listOf(
|
||||||
|
"Стой. — сказал он.",
|
||||||
|
"Стой! — сказал он.",
|
||||||
|
"Стой? — спросил он.",
|
||||||
|
"«Стой!» — сказал он.",
|
||||||
|
"Стой? - спросил он."
|
||||||
|
).forEach { text ->
|
||||||
|
assertEquals(listOf(text), RussianSpeechTextSplitter.splitSentences(text))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test
|
||||||
|
fun stillSeparatesOrdinarySentences() {
|
||||||
|
assertEquals(
|
||||||
|
listOf("Первое предложение.", "Второе предложение!", "Третье?"),
|
||||||
|
RussianSpeechTextSplitter.splitSentences("Первое предложение. Второе предложение! Третье?")
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -38,6 +38,10 @@ Before accepting the model, synthesize a held-out Russian validation list at mul
|
|||||||
`noise_scale`, and `noise_scale_w` settings, measure real-time factor on the target phone, and listen for skipped
|
`noise_scale`, and `noise_scale_w` settings, measure real-time factor on the target phone, and listen for skipped
|
||||||
words, unstable stress, clicks, and repeated phonemes. A successful export alone is not a quality gate.
|
words, unstable stress, clicks, and repeated phonemes. A successful export alone is not a quality gate.
|
||||||
|
|
||||||
|
The epoch-500 book-reading profile selected by listening comparison is stored in
|
||||||
|
`aletheia_book_profile.json`. Keep its duration, noise, and sentence-pause values together when evaluating the ONNX
|
||||||
|
model or integrating it into the reader.
|
||||||
|
|
||||||
On native Windows, first run the one-batch CUDA check:
|
On native Windows, first run the one-batch CUDA check:
|
||||||
|
|
||||||
```powershell
|
```powershell
|
||||||
|
|||||||
@@ -33,3 +33,53 @@ before producing WAV files.
|
|||||||
`train_piper_windows.ps1` provides the equivalent native Windows path and a `-SmokeTest` mode that executes one
|
`train_piper_windows.ps1` provides the equivalent native Windows path and a `-SmokeTest` mode that executes one
|
||||||
training and validation batch on CUDA before a long run. See `GPU_TRAINING.md` for host prerequisites and the
|
training and validation batch on CUDA before a long run. See `GPU_TRAINING.md` for host prerequisites and the
|
||||||
acceptance gate.
|
acceptance gate.
|
||||||
|
|
||||||
|
## Vosk BookTTS compression experiment
|
||||||
|
|
||||||
|
`vosk_booktts_experiment.py` reproduces the Android Vosk frontend, prepares real acoustic-model feeds from
|
||||||
|
`booktts_calibration_ru.txt`, performs calibrated QDQ int8 quantization of convolution layers, and renders
|
||||||
|
deterministic baseline/candidate WAV pairs. Keep generated `.npz`, candidate ONNX, and WAV files outside tracked
|
||||||
|
source directories, for example under `.codex-temp`. Do not replace the application asset until the candidate
|
||||||
|
loads with ONNX Runtime, runs on the target phone, and passes listening comparison.
|
||||||
|
Use `prepare --bert-model` plus `compare-feeds` to isolate a BERT replacement while keeping the acoustic model
|
||||||
|
identical in both WAV branches. `compare-feeds --resume` preserves existing WAVs and atomically updates its
|
||||||
|
report after every pair; `run_booktts_ab_windows.ps1` is the persistent Windows entry point.
|
||||||
|
|
||||||
|
`booktts_student.py` builds token-level teacher targets from the current int8 BERT and trains a four-layer,
|
||||||
|
256-dimensional Transformer student with a 768-dimensional drop-in output. Its small calibration corpus is only
|
||||||
|
for validating the training/export pipeline; a candidate for the application requires a much larger licensed
|
||||||
|
literary corpus, held-out evaluation, ONNX quantization, and listening tests. The optional
|
||||||
|
`--max-training-batches` and `--max-validation-batches` limits are for CUDA memory smoke tests only.
|
||||||
|
Each completed epoch atomically replaces `checkpoint.latest.pt`; pass `--resume` with the same output directory
|
||||||
|
to continue an interrupted run.
|
||||||
|
|
||||||
|
`run_booktts_student_windows.ps1` is the native Windows CUDA entry point for a full or resumed student run.
|
||||||
|
The exported student includes the dataset's original-to-compact token map so its `input_ids` remain compatible
|
||||||
|
with the existing Vosk vocabulary. `patch_booktts_student_token_map.py` adds the same lookup to an older export
|
||||||
|
without retraining.
|
||||||
|
|
||||||
|
`compare_booktts_bert.py` measures teacher/student and FP32/INT8 embedding error, inference time, and unseen-token
|
||||||
|
coverage on a held-out text file.
|
||||||
|
|
||||||
|
`build_booktts_prosody_corpus.py` creates a deterministic balance of questions, exclamations, quotations,
|
||||||
|
ellipsis, and neutral prose. Build its teacher dataset with `--reuse-token-map`, then fine-tune from the prior
|
||||||
|
`student.pt` with `--initial-model`; the held-out listening phrases must be passed through `--exclude`.
|
||||||
|
`run_booktts_prosody_dataset_windows.ps1` builds this reused-vocabulary teacher dataset persistently on Windows.
|
||||||
|
After auditing that dataset, `run_booktts_prosody_finetune_windows.ps1` starts the low-learning-rate CUDA run.
|
||||||
|
|
||||||
|
`extract_wikisource_corpus.py` streams an official Russian Wikisource XML/BZip2 dump and writes a deduplicated,
|
||||||
|
filtered sentence corpus plus provenance metadata. The source URL and dump size must be retained. Wikisource is
|
||||||
|
not a blanket rights clearance for every included work, so review the resulting corpus and applicable source
|
||||||
|
licenses before distributing a trained model.
|
||||||
|
|
||||||
|
`run_booktts_corpus_windows.ps1` completes the resumable BITS download on the CUDA host, verifies the official
|
||||||
|
dump size and SHA-1, extracts the configured number of sentences, and builds the streamed float16 teacher
|
||||||
|
dataset. It intentionally stops before training so vocabulary coverage and a held-out evaluation plan can be
|
||||||
|
reviewed first.
|
||||||
|
|
||||||
|
`audit_booktts_text_corpus.py` records the corpus SHA-256, sentence-length and word-count distributions,
|
||||||
|
residual wiki-markup counters, and deterministic review samples. Run it before training; its report is a
|
||||||
|
technical quality check and does not replace a rights review of the source works.
|
||||||
|
|
||||||
|
`audit_booktts_dataset.py` opens every generated teacher shard and verifies sample counts, tensor shapes,
|
||||||
|
token ranges, binary masks, and finite teacher values. It also prints a reproducible SHA-256 over shard bytes.
|
||||||
|
|||||||
@@ -0,0 +1,12 @@
|
|||||||
|
{
|
||||||
|
"id": "aletheia-book-moderate-v1",
|
||||||
|
"model_milestone": 500,
|
||||||
|
"length_scale": 1.5,
|
||||||
|
"noise_scale": 0.75,
|
||||||
|
"noise_w_scale": 0.95,
|
||||||
|
"sentence_silence_seconds": 0.25,
|
||||||
|
"selection": {
|
||||||
|
"method": "human_listening_test",
|
||||||
|
"result": "preferred_over_pause_only_and_stronger_variation"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,66 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Validate every shard of a BookTTS teacher dataset."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
|
||||||
|
def audit(args: argparse.Namespace) -> None:
|
||||||
|
metadata_path = args.dataset / "dataset.json"
|
||||||
|
metadata = json.loads(metadata_path.read_text(encoding="utf-8"))
|
||||||
|
token_map = np.load(args.dataset / "token_id_map.npy")
|
||||||
|
samples = 0
|
||||||
|
tokens = 0
|
||||||
|
bytes_total = metadata_path.stat().st_size + (args.dataset / "token_id_map.npy").stat().st_size
|
||||||
|
digest = hashlib.sha256()
|
||||||
|
for record in metadata.get("shards", []):
|
||||||
|
path = args.dataset / record["file"]
|
||||||
|
digest.update(path.read_bytes())
|
||||||
|
bytes_total += path.stat().st_size
|
||||||
|
with np.load(path) as shard:
|
||||||
|
ids = shard["original_input_ids"]
|
||||||
|
mask = shard["attention_mask"]
|
||||||
|
teacher = shard["teacher"]
|
||||||
|
if ids.shape != mask.shape or teacher.shape[:2] != ids.shape or teacher.shape[2:] != (768,):
|
||||||
|
raise ValueError(f"Incompatible shapes in {path.name}: {ids.shape}, {mask.shape}, {teacher.shape}")
|
||||||
|
if ids.shape[0] != record["samples"]:
|
||||||
|
raise ValueError(f"Sample count mismatch in {path.name}")
|
||||||
|
if ids.size and (ids.min() < 0 or ids.max() >= token_map.shape[0]):
|
||||||
|
raise ValueError(f"Token id outside token map in {path.name}")
|
||||||
|
if not np.isfinite(teacher).all():
|
||||||
|
raise ValueError(f"Non-finite teacher value in {path.name}")
|
||||||
|
if not np.all((mask == 0) | (mask == 1)):
|
||||||
|
raise ValueError(f"Non-binary attention mask in {path.name}")
|
||||||
|
samples += ids.shape[0]
|
||||||
|
tokens += int(mask.sum())
|
||||||
|
if samples != metadata["samples"]:
|
||||||
|
raise ValueError(f"Dataset sample count is {samples}, metadata says {metadata['samples']}")
|
||||||
|
report = {
|
||||||
|
"dataset": str(args.dataset),
|
||||||
|
"samples": samples,
|
||||||
|
"shards": len(metadata.get("shards", [])),
|
||||||
|
"tokens": tokens,
|
||||||
|
"vocab_size": metadata["vocab_size"],
|
||||||
|
"token_map_entries": int(token_map.shape[0]),
|
||||||
|
"bytes": bytes_total,
|
||||||
|
"shards_sha256": digest.hexdigest().upper(),
|
||||||
|
"teacher_dtype": metadata.get("teacher_dtype"),
|
||||||
|
}
|
||||||
|
print(json.dumps(report, ensure_ascii=False, indent=2))
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--dataset", type=Path, required=True)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(audit(parse_args()) or 0)
|
||||||
@@ -0,0 +1,87 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Audit a BookTTS text corpus and emit reproducible statistics and samples."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import random
|
||||||
|
import re
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
MARKUP_PATTERNS = ("{{", "}}", "[[", "]]", "<ref", "</", "{|", "|}", "http://", "https://")
|
||||||
|
RUSSIAN = re.compile(r"[А-Яа-яЁё]")
|
||||||
|
LETTERS = re.compile(r"[^\W\d_]", re.UNICODE)
|
||||||
|
|
||||||
|
|
||||||
|
def percentile(sorted_values: list[int], fraction: float) -> int:
|
||||||
|
if not sorted_values:
|
||||||
|
return 0
|
||||||
|
return sorted_values[round((len(sorted_values) - 1) * fraction)]
|
||||||
|
|
||||||
|
|
||||||
|
def audit(args: argparse.Namespace) -> None:
|
||||||
|
rng = random.Random(args.seed)
|
||||||
|
lengths: list[int] = []
|
||||||
|
words: list[int] = []
|
||||||
|
samples: list[str] = []
|
||||||
|
suspicious = {pattern: 0 for pattern in MARKUP_PATTERNS}
|
||||||
|
low_russian_ratio = 0
|
||||||
|
with_digits = 0
|
||||||
|
with_dialogue_dash = 0
|
||||||
|
digest = hashlib.sha256()
|
||||||
|
with args.corpus.open("rb") as binary:
|
||||||
|
for raw in binary:
|
||||||
|
digest.update(raw)
|
||||||
|
line = raw.decode("utf-8").rstrip("\r\n")
|
||||||
|
index = len(lengths)
|
||||||
|
lengths.append(len(line)); words.append(len(line.split()))
|
||||||
|
if any(char.isdigit() for char in line):
|
||||||
|
with_digits += 1
|
||||||
|
if line.startswith(("—", "–", "-")):
|
||||||
|
with_dialogue_dash += 1
|
||||||
|
letters = LETTERS.findall(line)
|
||||||
|
russian = RUSSIAN.findall(line)
|
||||||
|
if letters and len(russian) / len(letters) < 0.75:
|
||||||
|
low_russian_ratio += 1
|
||||||
|
for pattern in MARKUP_PATTERNS:
|
||||||
|
if pattern in line:
|
||||||
|
suspicious[pattern] += 1
|
||||||
|
if len(samples) < args.samples:
|
||||||
|
samples.append(line)
|
||||||
|
else:
|
||||||
|
replacement = rng.randint(0, index)
|
||||||
|
if replacement < args.samples:
|
||||||
|
samples[replacement] = line
|
||||||
|
ordered_lengths = sorted(lengths); ordered_words = sorted(words)
|
||||||
|
report = {
|
||||||
|
"corpus": str(args.corpus), "sha256": digest.hexdigest().upper(), "lines": len(lengths),
|
||||||
|
"bytes": args.corpus.stat().st_size,
|
||||||
|
"characters": {"min": min(lengths, default=0), "mean": sum(lengths) / len(lengths) if lengths else 0, "p50": percentile(ordered_lengths, 0.50), "p95": percentile(ordered_lengths, 0.95), "max": max(lengths, default=0)},
|
||||||
|
"words": {"min": min(words, default=0), "mean": sum(words) / len(words) if words else 0, "p50": percentile(ordered_words, 0.50), "p95": percentile(ordered_words, 0.95), "max": max(words, default=0)},
|
||||||
|
"with_digits": with_digits, "with_dialogue_dash": with_dialogue_dash,
|
||||||
|
"below_75_percent_russian_letters": low_russian_ratio, "suspicious_markup": suspicious,
|
||||||
|
"sample_seed": args.seed, "samples": samples,
|
||||||
|
}
|
||||||
|
rendered = json.dumps(report, ensure_ascii=False, indent=2) + "\n"
|
||||||
|
if args.output:
|
||||||
|
if args.output.exists():
|
||||||
|
raise FileExistsError(f"Refusing to overwrite {args.output}")
|
||||||
|
args.output.write_text(rendered, encoding="utf-8")
|
||||||
|
print(rendered, end="")
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--corpus", type=Path, required=True)
|
||||||
|
parser.add_argument("--output", type=Path)
|
||||||
|
parser.add_argument("--samples", type=int, default=20)
|
||||||
|
parser.add_argument("--seed", type=int, default=20260721)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(audit(parse_args()) or 0)
|
||||||
@@ -0,0 +1,100 @@
|
|||||||
|
— Что-то случилось? — тихо спросил он, остановившись у двери.
|
||||||
|
Кое-кто торопится, кое-что забывает, а кто-то всё-таки возвращается.
|
||||||
|
«Послушай, — сказала Ксения, — нам необходимо поговорить серьёзно».
|
||||||
|
Он открыл письмо и прочитал: решение принято; отступать уже поздно.
|
||||||
|
В комнате было тихо... Слишком тихо, чтобы не заметить чужого дыхания.
|
||||||
|
Почему вы молчите? Неужели ответ оказался настолько неожиданным?
|
||||||
|
Стой! Не делай ни шага, пока я не объясню происходящее.
|
||||||
|
На столе лежали старые часы (они давно остановились), ключ и записка.
|
||||||
|
Северо-западный ветер постепенно стих, но море оставалось беспокойным.
|
||||||
|
Едва заметная улыбка появилась на её лице и сразу исчезла.
|
||||||
|
Во-первых, нужно проверить дорогу; во-вторых, дождаться остальных.
|
||||||
|
Он говорил медленно, отчётливо и спокойно, не проглатывая окончания слов.
|
||||||
|
Слова «замок», «мука» и «атлас» меняют ударение вместе со значением.
|
||||||
|
У ворот стоял старый замок, а вдали виднелся заброшенный замок.
|
||||||
|
В 2026 году экспедиция прошла 12 345 километров за 21 день.
|
||||||
|
Температура опустилась до -12,05 градуса, а давление выросло на +7 единиц.
|
||||||
|
Поезд номер 5 отправляется в 07:30; посадка закончится через 10 минут.
|
||||||
|
Глава вторая. Ночная встреча у реки.
|
||||||
|
Господин Петров, ул. Лесная, дом 8 — так было написано на конверте.
|
||||||
|
Когда предложение переходит на следующую страницу, оно должно звучать непрерывно.
|
||||||
|
Рассказчик сделал короткую паузу — ровно настолько, чтобы сохранить напряжение.
|
||||||
|
После двоеточия пауза короче: мысль продолжается без завершения фразы.
|
||||||
|
После точки, вопросительного и восклицательного знаков нужна ясная граница.
|
||||||
|
Чёткая дикция важнее скорости, особенно в длинных книжных предложениях.
|
||||||
|
За окном, где ещё недавно шумел дождь, медленно светлело утреннее небо.
|
||||||
|
— Подождите меня здесь, — произнёс проводник, — и никуда не уходите.
|
||||||
|
Она хотела возразить, однако передумала: спор всё равно ничего бы не изменил.
|
||||||
|
Если дверь заперта, постучи трижды; если никто не ответит — возвращайся.
|
||||||
|
Где-то далеко прокричала ночная птица, и лес снова погрузился в тишину.
|
||||||
|
Как странно: знакомый дом казался теперь меньше, темнее и гораздо старше.
|
||||||
|
«Я вернусь до рассвета!» — крикнул он и исчез за поворотом.
|
||||||
|
Неужели всё это — лишь совпадение, которому мы напрасно придали значение?
|
||||||
|
Ветер перебирал сухие листья — шорох за шорохом, вздох за вздохом.
|
||||||
|
Она подняла глаза (в них блестели слёзы), но голос её остался твёрдым.
|
||||||
|
Старик помолчал; затем, словно решившись, заговорил совсем другим тоном.
|
||||||
|
Ни света, ни звука, ни единого следа — только холодная пустая дорога.
|
||||||
|
Прежде чем ответить, он медленно сложил письмо и убрал его во внутренний карман.
|
||||||
|
— Вы уверены? — Да. — Тогда начинайте, времени у нас почти не осталось.
|
||||||
|
По ту сторону реки поднимался древний лес, окутанный сизым предрассветным туманом.
|
||||||
|
Всё-таки, по-моему, кто-нибудь должен был предупредить нас заранее.
|
||||||
|
Из-за полуоткрытой двери донёсся чей-то приглушённый, едва различимый голос.
|
||||||
|
В письме стояло одно слово: «Жди»; ни подписи, ни даты не было.
|
||||||
|
Она замерла на полуслове, будто внезапно услышала то, чего не слышали остальные.
|
||||||
|
Гром прогремел совсем близко — окна задрожали, а свеча внезапно погасла.
|
||||||
|
Мальчик нёс в руках старинный атлас, а на плечах у него лежал атласный плащ.
|
||||||
|
Мука́ закончилась к полудню, но му́ка ожидания продолжалась до самого вечера.
|
||||||
|
На берегу белели старые и́рисы, а в письме лежал неоплаченный ири́с.
|
||||||
|
Острый клино́к сверкнул в темноте, и сухой кли́нок дерева хрустнул под ногой.
|
||||||
|
Мы обошли весь квартал, но нужный дом так и не нашли.
|
||||||
|
Он поставил подпись под договором, хотя каждый пункт вызывал у него сомнения.
|
||||||
|
В первой главе герой уезжает; во второй — возвращается под чужим именем.
|
||||||
|
Запомни главное: нельзя перебивать фразу только потому, что закончилась страница.
|
||||||
|
Последние слова предложения должны прозвучать на следующем экране без повторения начала.
|
||||||
|
Она читала неторопливо, оставляя между словами едва заметное пространство.
|
||||||
|
Слишком быстрый темп съедает согласные, а чрезмерно медленный разрушает смысл фразы.
|
||||||
|
На отметке 3,5 километра дорога раздваивается: налево — к озеру, направо — к селу.
|
||||||
|
В архиве значились тома № 7, 12 и 18; восьмого тома в описи не было.
|
||||||
|
Экспедиция началась 14 июля 1897 года и завершилась лишь через восемь месяцев.
|
||||||
|
В 6 часов 45 минут колокол ударил дважды, хотя должен был ударить шесть раз.
|
||||||
|
Расстояние составляло около 1 250 метров, то есть чуть больше одной версты.
|
||||||
|
Цена книги — 19 рублей 90 копеек; на обороте карандашом написано: «Не продавать».
|
||||||
|
Глава двенадцатая. Письмо, которого никто не ждал.
|
||||||
|
Часть III. Возвращение домой после двадцати лет странствий.
|
||||||
|
— Кто там? — Это я. — Кто «я»? — Откройте, и вы всё поймёте.
|
||||||
|
«Нет, нет и ещё раз нет», — отчётливо повторила она.
|
||||||
|
Он прошептал: «Тише… нас могут услышать», — и указал на окно.
|
||||||
|
Сначала послышались шаги; потом скрипнула лестница; наконец открылась дверь.
|
||||||
|
Дорога была трудна: снег слепил глаза, ветер сбивал с ног, силы иссякали.
|
||||||
|
Он не ответил — не потому, что не знал ответа, а потому, что боялся его произнести.
|
||||||
|
Кто бы мог подумать, что маленькая находка изменит судьбу целого города!
|
||||||
|
Когда часы пробили полночь, незнакомец снял шляпу и назвал своё настоящее имя.
|
||||||
|
Утром всё выглядело обыкновенно; только следы у калитки напоминали о ночном госте.
|
||||||
|
Я хотел было уйти, но она сказала: «Останьтесь ещё на одну минуту».
|
||||||
|
Голос звучал спокойно, почти равнодушно, и от этого становилось ещё тревожнее.
|
||||||
|
Вдали показалась станция — низкая платформа, жёлтый фонарь и одинокий смотритель.
|
||||||
|
Он читал старую рукопись строка за строкой, боясь пропустить хотя бы одну букву.
|
||||||
|
Некоторые слова были стёрты; другие — исправлены чужой рукой много лет спустя.
|
||||||
|
В примечании значилось: см. главу 4, стр. 217, абзац второй.
|
||||||
|
На табличке было написано: «Вход воспрещён с 22:00 до 06:00».
|
||||||
|
Температура воды равнялась 18,4 °C, скорость течения — 2,1 метра в секунду.
|
||||||
|
Вероятность ошибки оценили в 0,03 процента, но случай всё-таки произошёл.
|
||||||
|
Координаты точки: 55°45′ северной широты, 37°37′ восточной долготы.
|
||||||
|
В комнате находились А. П. Чехов, Л. Н. Толстой и ещё двое гостей.
|
||||||
|
Иван Сергеевич, будьте добры, прочтите последний абзац ещё раз.
|
||||||
|
Мать-и-мачеха росла у дороги, а где-то рядом кричала птица-пересмешник.
|
||||||
|
Поезд шёл с северо-востока на юго-запад, постепенно набирая скорость.
|
||||||
|
Что-либо менять было поздно; кое-как собрав вещи, они отправились в путь.
|
||||||
|
Кто-нибудь видел серо-зелёную папку, лежавшую здесь полчаса назад?
|
||||||
|
Из-под земли доносился глухой рокот, то усиливаясь, то почти исчезая.
|
||||||
|
Он всё повторял одно и то же, словно заученное заклинание: «Нельзя опаздывать».
|
||||||
|
Свет погас на мгновение — и именно в это мгновение картина исчезла со стены.
|
||||||
|
Слово за словом, страница за страницей перед читателем раскрывалась чужая жизнь.
|
||||||
|
Её ответ был прост, но окончателен: она не вернётся ни завтра, ни через год.
|
||||||
|
Вопрос состоял не в том, кто виноват, а в том, можно ли ещё что-нибудь исправить.
|
||||||
|
Не открывая глаз, он прислушался: рядом потрескивал огонь, за стеной шумела вода.
|
||||||
|
«Вы опоздали на семь минут», — заметил человек в сером пальто.
|
||||||
|
На мгновение ей показалось, будто портрет улыбнулся; разумеется, этого не могло быть.
|
||||||
|
Туман рассеивался медленно, открывая то крышу, то колокольню, то дальний берег.
|
||||||
|
Никто не произнёс ни слова — прощание и без того затянулось.
|
||||||
|
И всё же где-то в глубине души оставалась слабая, почти невозможная надежда.
|
||||||
@@ -0,0 +1,4 @@
|
|||||||
|
Почему вы молчите? Неужели ответ оказался настолько неожиданным?
|
||||||
|
В 2026 году экспедиция прошла 12 345 километров за 21 день.
|
||||||
|
Мука́ закончилась к полудню, но му́ка ожидания продолжалась до самого вечера.
|
||||||
|
Мать-и-мачеха росла у дороги, а где-то рядом кричала птица-пересмешник.
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
Мука́ закончилась к полудню, но му́ка ожидания продолжалась до самого вечера.
|
||||||
@@ -0,0 +1,5 @@
|
|||||||
|
Почему вы молчите? Неужели ответ оказался настолько неожиданным?
|
||||||
|
— Что-то случилось? — тихо спросил он, остановившись у двери.
|
||||||
|
Чёткая дикция важнее скорости, особенно в длинных книжных предложениях.
|
||||||
|
Неужели всё это — лишь совпадение, которому мы напрасно придали значение?
|
||||||
|
И всё же где-то в глубине души оставалась слабая, почти невозможная надежда.
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
— Что-то случилось?
|
||||||
|
Почему вы молчите?
|
||||||
|
Неужели ответ оказался настолько неожиданным?
|
||||||
@@ -0,0 +1,435 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Build and train a compact drop-in BERT student for Aletheia BookTTS."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import random
|
||||||
|
import shutil
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
from collections import Counter
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from vosk_booktts_experiment import WordPieceTokenizer, load_corpus, normalize_text
|
||||||
|
|
||||||
|
|
||||||
|
def build_dataset(args: argparse.Namespace) -> None:
|
||||||
|
import onnxruntime as ort
|
||||||
|
|
||||||
|
args.output.mkdir(parents=True, exist_ok=False)
|
||||||
|
tokenizer = WordPieceTokenizer(args.assets / "vocab.txt")
|
||||||
|
options = ort.SessionOptions()
|
||||||
|
options.log_severity_level = 3
|
||||||
|
teacher = ort.InferenceSession(
|
||||||
|
str(args.assets / "bert.int8.onnx"),
|
||||||
|
sess_options=options,
|
||||||
|
providers=["CPUExecutionProvider"],
|
||||||
|
)
|
||||||
|
token_counts: Counter[int] = Counter()
|
||||||
|
pending: list[tuple[str, np.ndarray]] = []
|
||||||
|
shard_samples: list[tuple[np.ndarray, np.ndarray]] = []
|
||||||
|
shard_records = []
|
||||||
|
sample_index = 0
|
||||||
|
shard_index = 0
|
||||||
|
teacher_dtype = np.float16 if args.teacher_dtype == "float16" else np.float32
|
||||||
|
|
||||||
|
def flush_shard() -> None:
|
||||||
|
nonlocal shard_index
|
||||||
|
if not shard_samples:
|
||||||
|
return
|
||||||
|
longest = max(ids.size for ids, _ in shard_samples)
|
||||||
|
shard_ids = np.zeros((len(shard_samples), longest), dtype=np.int64)
|
||||||
|
shard_mask = np.zeros_like(shard_ids)
|
||||||
|
shard_teacher = np.zeros((len(shard_samples), longest, 768), dtype=teacher_dtype)
|
||||||
|
for row, (ids, embeddings) in enumerate(shard_samples):
|
||||||
|
shard_ids[row, : ids.size] = ids
|
||||||
|
shard_mask[row, : ids.size] = 1
|
||||||
|
shard_teacher[row, : ids.size] = embeddings
|
||||||
|
filename = f"shard-{shard_index:05d}.npz"
|
||||||
|
np.savez(
|
||||||
|
args.output / filename,
|
||||||
|
original_input_ids=shard_ids,
|
||||||
|
attention_mask=shard_mask,
|
||||||
|
teacher=shard_teacher,
|
||||||
|
)
|
||||||
|
shard_records.append({"file": filename, "samples": len(shard_samples), "max_tokens": longest})
|
||||||
|
shard_samples.clear()
|
||||||
|
shard_index += 1
|
||||||
|
|
||||||
|
def flush_batch(manifest) -> None:
|
||||||
|
nonlocal sample_index
|
||||||
|
if not pending:
|
||||||
|
return
|
||||||
|
longest = max(ids.size for _, ids in pending)
|
||||||
|
batch_ids = np.zeros((len(pending), longest), dtype=np.int64)
|
||||||
|
batch_mask = np.zeros_like(batch_ids)
|
||||||
|
for row, (_, ids) in enumerate(pending):
|
||||||
|
batch_ids[row, : ids.size] = ids
|
||||||
|
batch_mask[row, : ids.size] = 1
|
||||||
|
outputs = teacher.run(None, {
|
||||||
|
"input_ids": batch_ids,
|
||||||
|
"attention_mask": batch_mask,
|
||||||
|
"token_type_ids": np.zeros_like(batch_ids),
|
||||||
|
})[0]
|
||||||
|
if outputs.ndim == 2:
|
||||||
|
outputs = outputs[np.newaxis, :, :]
|
||||||
|
for row, (text, ids) in enumerate(pending):
|
||||||
|
embeddings = outputs[row, : ids.size].astype(teacher_dtype, copy=False)
|
||||||
|
if embeddings.shape != (ids.size, 768):
|
||||||
|
raise ValueError(f"Unexpected teacher shape {embeddings.shape}")
|
||||||
|
shard_samples.append((ids.copy(), embeddings.copy()))
|
||||||
|
manifest.write(json.dumps({"sample": sample_index, "text": text, "tokens": int(ids.size)}, ensure_ascii=False) + "\n")
|
||||||
|
sample_index += 1
|
||||||
|
if len(shard_samples) >= args.shard_size:
|
||||||
|
flush_shard()
|
||||||
|
pending.clear()
|
||||||
|
|
||||||
|
with (args.output / "manifest.jsonl").open("w", encoding="utf-8", newline="\n") as manifest:
|
||||||
|
for text in load_corpus(args.corpus):
|
||||||
|
ids, _ = tokenizer.encode(normalize_text(text))
|
||||||
|
token_counts.update(int(value) for value in ids)
|
||||||
|
pending.append((text, ids))
|
||||||
|
if len(pending) >= args.teacher_batch_size:
|
||||||
|
flush_batch(manifest)
|
||||||
|
if sample_index % 1_000 == 0:
|
||||||
|
print(f"teacher_samples={sample_index}", flush=True)
|
||||||
|
flush_batch(manifest)
|
||||||
|
flush_shard()
|
||||||
|
|
||||||
|
observed = set(token_counts)
|
||||||
|
if args.reuse_token_map:
|
||||||
|
source_metadata = json.loads((args.reuse_token_map / "dataset.json").read_text(encoding="utf-8"))
|
||||||
|
source_map = np.load(args.reuse_token_map / "token_id_map.npy")
|
||||||
|
if source_map.size != len(tokenizer.tokens):
|
||||||
|
raise ValueError("Reused token map and teacher vocabulary have different sizes")
|
||||||
|
shutil.copy2(args.reuse_token_map / "token_id_map.npy", args.output / "token_id_map.npy")
|
||||||
|
shutil.copy2(args.reuse_token_map / "vocab.txt", args.output / "vocab.txt")
|
||||||
|
vocab_size = int(source_metadata["vocab_size"])
|
||||||
|
unknown_new = int(source_map[tokenizer.vocabulary["[UNK]"]])
|
||||||
|
replaced_occurrences = sum(
|
||||||
|
count for token, count in token_counts.items()
|
||||||
|
if token != tokenizer.vocabulary["[UNK]"] and int(source_map[token]) == unknown_new
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
special_ids = [tokenizer.vocabulary[name] for name in ("[PAD]", "[UNK]", "[CLS]", "[SEP]")]
|
||||||
|
if args.max_vocab and len(observed) > args.max_vocab:
|
||||||
|
retained = set(special_ids)
|
||||||
|
retained.update(token for token, _ in token_counts.most_common(args.max_vocab - len(retained)))
|
||||||
|
else:
|
||||||
|
retained = observed | set(special_ids)
|
||||||
|
ordered_old_ids = special_ids + sorted(retained - set(special_ids))
|
||||||
|
old_to_new = {old: new for new, old in enumerate(ordered_old_ids)}
|
||||||
|
unknown_new = old_to_new[tokenizer.vocabulary["[UNK]"]]
|
||||||
|
reduced_tokens = [tokenizer.tokens[index] for index in ordered_old_ids]
|
||||||
|
(args.output / "vocab.txt").write_text("\n".join(reduced_tokens) + "\n", encoding="utf-8")
|
||||||
|
token_id_map = np.full(len(tokenizer.tokens), unknown_new, dtype=np.int64)
|
||||||
|
for old, new in old_to_new.items():
|
||||||
|
token_id_map[old] = new
|
||||||
|
np.save(args.output / "token_id_map.npy", token_id_map)
|
||||||
|
vocab_size = len(reduced_tokens)
|
||||||
|
replaced_occurrences = sum(count for token, count in token_counts.items() if token not in retained)
|
||||||
|
metadata = {
|
||||||
|
"samples": sample_index,
|
||||||
|
"vocab_size": vocab_size,
|
||||||
|
"observed_original_tokens": len(observed),
|
||||||
|
"replaced_token_occurrences": replaced_occurrences,
|
||||||
|
"reused_token_map": str(args.reuse_token_map) if args.reuse_token_map else None,
|
||||||
|
"teacher_dtype": args.teacher_dtype,
|
||||||
|
"teacher_batch_size": args.teacher_batch_size,
|
||||||
|
"shard_size": args.shard_size,
|
||||||
|
"shards": shard_records,
|
||||||
|
}
|
||||||
|
(args.output / "dataset.json").write_text(json.dumps(metadata, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||||
|
print(json.dumps(metadata, indent=2))
|
||||||
|
|
||||||
|
|
||||||
|
def train_student(args: argparse.Namespace) -> None:
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
from torch.utils.data import DataLoader, Dataset, IterableDataset, Subset
|
||||||
|
|
||||||
|
class DistillationDataset(Dataset):
|
||||||
|
def __init__(self, root: Path):
|
||||||
|
self.files = sorted(root.glob("sample-*.npz"))
|
||||||
|
if not self.files:
|
||||||
|
raise ValueError(f"No samples in {root}")
|
||||||
|
mapping = root / "token_id_map.npy"
|
||||||
|
self.token_id_map = np.load(mapping) if mapping.exists() else None
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return len(self.files)
|
||||||
|
|
||||||
|
def __getitem__(self, index: int) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
with np.load(self.files[index]) as data:
|
||||||
|
if "original_input_ids" in data:
|
||||||
|
original = data["original_input_ids"]
|
||||||
|
if self.token_id_map is None:
|
||||||
|
raise ValueError("token_id_map.npy is required for original_input_ids")
|
||||||
|
input_ids = self.token_id_map[original]
|
||||||
|
else:
|
||||||
|
input_ids = data["input_ids"]
|
||||||
|
teacher = data["teacher"].astype(np.float32)
|
||||||
|
return torch.from_numpy(input_ids.copy()), torch.from_numpy(teacher.copy())
|
||||||
|
|
||||||
|
class ShardedDataset(IterableDataset):
|
||||||
|
def __init__(self, root: Path, records: list[dict], shuffle: bool, seed: int):
|
||||||
|
self.root = root
|
||||||
|
self.records = list(records)
|
||||||
|
self.shuffle = shuffle
|
||||||
|
self.seed = seed
|
||||||
|
self.epoch = 0
|
||||||
|
self.token_id_map = np.load(root / "token_id_map.npy")
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return sum(record["samples"] for record in self.records)
|
||||||
|
|
||||||
|
def set_epoch(self, epoch: int) -> None:
|
||||||
|
self.epoch = epoch
|
||||||
|
|
||||||
|
def __iter__(self):
|
||||||
|
rng = random.Random(self.seed + self.epoch)
|
||||||
|
records = list(self.records)
|
||||||
|
if self.shuffle:
|
||||||
|
rng.shuffle(records)
|
||||||
|
for record in records:
|
||||||
|
with np.load(self.root / record["file"]) as data:
|
||||||
|
original = data["original_input_ids"]
|
||||||
|
masks = data["attention_mask"]
|
||||||
|
teachers = data["teacher"]
|
||||||
|
rows = list(range(original.shape[0]))
|
||||||
|
if self.shuffle:
|
||||||
|
rng.shuffle(rows)
|
||||||
|
for row in rows:
|
||||||
|
length = int(masks[row].sum())
|
||||||
|
ids = self.token_id_map[original[row, :length]]
|
||||||
|
teacher = teachers[row, :length].astype(np.float32)
|
||||||
|
yield torch.from_numpy(ids.copy()), torch.from_numpy(teacher.copy())
|
||||||
|
|
||||||
|
def collate(batch):
|
||||||
|
longest = max(ids.size(0) for ids, _ in batch)
|
||||||
|
ids = torch.zeros((len(batch), longest), dtype=torch.long)
|
||||||
|
mask = torch.zeros((len(batch), longest), dtype=torch.long)
|
||||||
|
targets = torch.zeros((len(batch), longest, 768), dtype=torch.float32)
|
||||||
|
for row, (sample_ids, teacher) in enumerate(batch):
|
||||||
|
length = sample_ids.size(0)
|
||||||
|
ids[row, :length] = sample_ids
|
||||||
|
mask[row, :length] = 1
|
||||||
|
targets[row, :length] = teacher
|
||||||
|
return ids, mask, targets
|
||||||
|
|
||||||
|
class Student(nn.Module):
|
||||||
|
def __init__(self, vocab_size: int):
|
||||||
|
super().__init__()
|
||||||
|
self.token_embedding = nn.Embedding(vocab_size, args.hidden, padding_idx=0)
|
||||||
|
self.position_embedding = nn.Embedding(args.max_length, args.hidden)
|
||||||
|
self.type_embedding = nn.Embedding(2, args.hidden)
|
||||||
|
layer = nn.TransformerEncoderLayer(
|
||||||
|
d_model=args.hidden, nhead=args.heads, dim_feedforward=args.feed_forward,
|
||||||
|
dropout=args.dropout, activation="gelu", batch_first=True, norm_first=True,
|
||||||
|
)
|
||||||
|
self.encoder = nn.TransformerEncoder(layer, num_layers=args.layers, enable_nested_tensor=False)
|
||||||
|
self.norm = nn.LayerNorm(args.hidden)
|
||||||
|
self.projection = nn.Linear(args.hidden, 768)
|
||||||
|
|
||||||
|
def forward(self, input_ids, attention_mask, token_type_ids=None):
|
||||||
|
if token_type_ids is None:
|
||||||
|
token_type_ids = torch.zeros_like(input_ids)
|
||||||
|
positions = torch.arange(input_ids.size(1), device=input_ids.device).unsqueeze(0)
|
||||||
|
values = (
|
||||||
|
self.token_embedding(input_ids)
|
||||||
|
+ self.position_embedding(positions)
|
||||||
|
+ self.type_embedding(token_type_ids)
|
||||||
|
)
|
||||||
|
values = self.encoder(values, src_key_padding_mask=attention_mask == 0)
|
||||||
|
return self.projection(self.norm(values))
|
||||||
|
|
||||||
|
class ExportWrapper(nn.Module):
|
||||||
|
def __init__(self, model, token_id_map):
|
||||||
|
super().__init__(); self.model = model
|
||||||
|
self.register_buffer("token_id_map", token_id_map)
|
||||||
|
|
||||||
|
def forward(self, input_ids, attention_mask, token_type_ids):
|
||||||
|
mapped_input_ids = self.token_id_map[input_ids]
|
||||||
|
return self.model(mapped_input_ids, attention_mask, token_type_ids)[0]
|
||||||
|
|
||||||
|
random.seed(args.seed); np.random.seed(args.seed); torch.manual_seed(args.seed)
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.manual_seed_all(args.seed)
|
||||||
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
|
metadata = json.loads((args.dataset / "dataset.json").read_text(encoding="utf-8"))
|
||||||
|
model = Student(metadata["vocab_size"]).to(device)
|
||||||
|
if args.initial_model:
|
||||||
|
payload = torch.load(args.initial_model, map_location=device, weights_only=False)
|
||||||
|
state_dict = payload.get("state_dict") or payload.get("best_state") or payload.get("model")
|
||||||
|
if state_dict is None:
|
||||||
|
raise ValueError(f"No model state in {args.initial_model}")
|
||||||
|
model.load_state_dict(state_dict)
|
||||||
|
if metadata.get("shards"):
|
||||||
|
records = list(metadata["shards"])
|
||||||
|
rng = random.Random(args.seed); rng.shuffle(records)
|
||||||
|
validation_shards = max(1, round(len(records) * args.validation_fraction)) if len(records) > 1 else 0
|
||||||
|
validation_records = records[:validation_shards]
|
||||||
|
training_records = records[validation_shards:] or records
|
||||||
|
training_data = ShardedDataset(args.dataset, training_records, True, args.seed)
|
||||||
|
validation_data = ShardedDataset(args.dataset, validation_records, False, args.seed) if validation_records else None
|
||||||
|
loader = DataLoader(training_data, batch_size=args.batch_size, collate_fn=collate)
|
||||||
|
validation_loader = DataLoader(validation_data, batch_size=args.batch_size, collate_fn=collate) if validation_data else None
|
||||||
|
else:
|
||||||
|
all_data = DistillationDataset(args.dataset)
|
||||||
|
validation_size = max(1, round(len(all_data) * args.validation_fraction)) if len(all_data) > 1 else 0
|
||||||
|
indices = list(range(len(all_data))); random.Random(args.seed).shuffle(indices)
|
||||||
|
validation_data = Subset(all_data, indices[:validation_size]) if validation_size else None
|
||||||
|
training_data = Subset(all_data, indices[validation_size:] or indices)
|
||||||
|
loader = DataLoader(training_data, batch_size=args.batch_size, shuffle=True, collate_fn=collate)
|
||||||
|
validation_loader = DataLoader(validation_data, batch_size=args.batch_size, collate_fn=collate) if validation_data else None
|
||||||
|
optimizer = torch.optim.AdamW(model.parameters(), lr=args.learning_rate, weight_decay=0.01)
|
||||||
|
scaler = torch.amp.GradScaler("cuda", enabled=device.type == "cuda")
|
||||||
|
checkpoint_path = args.output / "checkpoint.latest.pt"
|
||||||
|
if args.resume:
|
||||||
|
if not checkpoint_path.exists():
|
||||||
|
raise FileNotFoundError(f"Cannot resume without {checkpoint_path}")
|
||||||
|
else:
|
||||||
|
args.output.mkdir(parents=True, exist_ok=False)
|
||||||
|
started = time.time(); history = []; best_loss = float("inf"); best_state = None; best_epoch = 0; stale_epochs = 0
|
||||||
|
start_epoch = 1
|
||||||
|
if args.resume:
|
||||||
|
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||||
|
model.load_state_dict(checkpoint["model"])
|
||||||
|
optimizer.load_state_dict(checkpoint["optimizer"])
|
||||||
|
scaler.load_state_dict(checkpoint["scaler"])
|
||||||
|
history = checkpoint["history"]
|
||||||
|
best_loss = checkpoint["best_loss"]
|
||||||
|
best_state = checkpoint["best_state"]
|
||||||
|
best_epoch = checkpoint["best_epoch"]
|
||||||
|
stale_epochs = checkpoint["stale_epochs"]
|
||||||
|
start_epoch = checkpoint["epoch"] + 1
|
||||||
|
|
||||||
|
def calculate_loss(prediction, target, mask):
|
||||||
|
active = mask.bool().unsqueeze(-1).expand_as(prediction)
|
||||||
|
mse = torch.mean((prediction[active] - target[active]) ** 2)
|
||||||
|
cosine = 1.0 - torch.nn.functional.cosine_similarity(prediction[mask.bool()], target[mask.bool()], dim=-1).mean()
|
||||||
|
return mse + args.cosine_weight * cosine
|
||||||
|
|
||||||
|
for epoch in range(start_epoch, args.epochs + 1):
|
||||||
|
if isinstance(training_data, ShardedDataset):
|
||||||
|
training_data.set_epoch(epoch)
|
||||||
|
model.train(); total = 0.0; batches = 0
|
||||||
|
for ids, mask, target in loader:
|
||||||
|
ids, mask, target = ids.to(device), mask.to(device), target.to(device)
|
||||||
|
optimizer.zero_grad(set_to_none=True)
|
||||||
|
with torch.amp.autocast("cuda", dtype=torch.float16, enabled=device.type == "cuda"):
|
||||||
|
prediction = model(ids, mask)
|
||||||
|
loss = calculate_loss(prediction, target, mask)
|
||||||
|
scaler.scale(loss).backward(); scaler.unscale_(optimizer)
|
||||||
|
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
||||||
|
scaler.step(optimizer); scaler.update()
|
||||||
|
total += float(loss.detach()); batches += 1
|
||||||
|
if args.max_training_batches and batches >= args.max_training_batches:
|
||||||
|
break
|
||||||
|
training_loss = total / batches
|
||||||
|
validation_loss = training_loss
|
||||||
|
if validation_loader is not None:
|
||||||
|
model.eval(); validation_total = 0.0; validation_batches = 0
|
||||||
|
with torch.no_grad():
|
||||||
|
for ids, mask, target in validation_loader:
|
||||||
|
ids, mask, target = ids.to(device), mask.to(device), target.to(device)
|
||||||
|
with torch.amp.autocast("cuda", dtype=torch.float16, enabled=device.type == "cuda"):
|
||||||
|
validation_total += float(calculate_loss(model(ids, mask), target, mask))
|
||||||
|
validation_batches += 1
|
||||||
|
if args.max_validation_batches and validation_batches >= args.max_validation_batches:
|
||||||
|
break
|
||||||
|
validation_loss = validation_total / validation_batches
|
||||||
|
history.append({"epoch": epoch, "training_loss": training_loss, "validation_loss": validation_loss})
|
||||||
|
print(f"epoch={epoch} training_loss={training_loss:.8f} validation_loss={validation_loss:.8f}")
|
||||||
|
if validation_loss < best_loss:
|
||||||
|
best_loss = validation_loss; best_epoch = epoch; stale_epochs = 0
|
||||||
|
best_state = {name: value.detach().cpu().clone() for name, value in model.state_dict().items()}
|
||||||
|
else:
|
||||||
|
stale_epochs += 1
|
||||||
|
checkpoint = {
|
||||||
|
"epoch": epoch, "model": model.state_dict(), "optimizer": optimizer.state_dict(),
|
||||||
|
"scaler": scaler.state_dict(), "history": history, "best_loss": best_loss,
|
||||||
|
"best_state": best_state, "best_epoch": best_epoch, "stale_epochs": stale_epochs,
|
||||||
|
}
|
||||||
|
temporary_checkpoint = checkpoint_path.with_suffix(".tmp")
|
||||||
|
torch.save(checkpoint, temporary_checkpoint); temporary_checkpoint.replace(checkpoint_path)
|
||||||
|
if stale_epochs >= args.early_stopping_patience:
|
||||||
|
print(f"early_stop epoch={epoch} best_epoch={best_epoch}")
|
||||||
|
break
|
||||||
|
|
||||||
|
if best_state is not None:
|
||||||
|
model.load_state_dict(best_state)
|
||||||
|
config = {key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items() if key != "handler"}
|
||||||
|
torch.save({"state_dict": model.state_dict(), "metadata": metadata, "config": config}, args.output / "student.pt")
|
||||||
|
model.eval()
|
||||||
|
export_token_map = torch.from_numpy(np.load(args.dataset / "token_id_map.npy")).long()
|
||||||
|
wrapper = ExportWrapper(model, export_token_map).cpu().eval()
|
||||||
|
example_ids = torch.tensor([[2, 3]], dtype=torch.long)
|
||||||
|
example_mask = torch.ones_like(example_ids); example_types = torch.zeros_like(example_ids)
|
||||||
|
sequence = torch.export.Dim("sequence", min=2, max=args.max_length)
|
||||||
|
torch.onnx.export(
|
||||||
|
wrapper, (example_ids, example_mask, example_types), args.output / "bert.student.fp32.onnx",
|
||||||
|
input_names=["input_ids", "attention_mask", "token_type_ids"], output_names=["logits"],
|
||||||
|
dynamic_shapes=({1: sequence}, {1: sequence}, {1: sequence}),
|
||||||
|
opset_version=18, dynamo=True, external_data=False,
|
||||||
|
)
|
||||||
|
report = {
|
||||||
|
"device": str(device), "parameters": sum(parameter.numel() for parameter in model.parameters()),
|
||||||
|
"training_samples": len(training_data), "validation_samples": len(validation_data) if validation_data is not None else 0,
|
||||||
|
"best_epoch": best_epoch, "best_validation_loss": best_loss,
|
||||||
|
"peak_cuda_bytes": torch.cuda.max_memory_allocated() if device.type == "cuda" else 0,
|
||||||
|
"seconds": time.time() - started, "onnx_bytes": (args.output / "bert.student.fp32.onnx").stat().st_size,
|
||||||
|
"history": history, "training_batches_last_epoch": batches,
|
||||||
|
"validation_batches_last_epoch": validation_batches if validation_loader is not None else 0,
|
||||||
|
}
|
||||||
|
(args.output / "training.json").write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8")
|
||||||
|
print(json.dumps({key: value for key, value in report.items() if key != "history"}, indent=2))
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
commands = parser.add_subparsers(dest="command", required=True)
|
||||||
|
dataset = commands.add_parser("dataset")
|
||||||
|
dataset.add_argument("--assets", type=Path, required=True)
|
||||||
|
dataset.add_argument("--corpus", type=Path, required=True)
|
||||||
|
dataset.add_argument("--output", type=Path, required=True)
|
||||||
|
dataset.add_argument("--max-vocab", type=int, default=0)
|
||||||
|
dataset.add_argument("--teacher-batch-size", type=int, default=16)
|
||||||
|
dataset.add_argument("--teacher-dtype", choices=("float16", "float32"), default="float16")
|
||||||
|
dataset.add_argument("--shard-size", type=int, default=256)
|
||||||
|
dataset.add_argument("--reuse-token-map", type=Path)
|
||||||
|
dataset.set_defaults(handler=build_dataset)
|
||||||
|
train = commands.add_parser("train")
|
||||||
|
train.add_argument("--dataset", type=Path, required=True)
|
||||||
|
train.add_argument("--output", type=Path, required=True)
|
||||||
|
train.add_argument("--epochs", type=int, default=10)
|
||||||
|
train.add_argument("--batch-size", type=int, default=8)
|
||||||
|
train.add_argument("--hidden", type=int, default=256)
|
||||||
|
train.add_argument("--heads", type=int, default=8)
|
||||||
|
train.add_argument("--feed-forward", type=int, default=768)
|
||||||
|
train.add_argument("--layers", type=int, default=4)
|
||||||
|
train.add_argument("--max-length", type=int, default=256)
|
||||||
|
train.add_argument("--dropout", type=float, default=0.1)
|
||||||
|
train.add_argument("--learning-rate", type=float, default=3e-4)
|
||||||
|
train.add_argument("--cosine-weight", type=float, default=0.1)
|
||||||
|
train.add_argument("--validation-fraction", type=float, default=0.05)
|
||||||
|
train.add_argument("--early-stopping-patience", type=int, default=5)
|
||||||
|
train.add_argument("--max-training-batches", type=int, default=0)
|
||||||
|
train.add_argument("--max-validation-batches", type=int, default=0)
|
||||||
|
train.add_argument("--resume", action="store_true")
|
||||||
|
train.add_argument("--initial-model", type=Path)
|
||||||
|
train.add_argument("--seed", type=int, default=20260721)
|
||||||
|
train.set_defaults(handler=train_student)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> int:
|
||||||
|
args = parse_args(); args.handler(args); return 0
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(main())
|
||||||
@@ -0,0 +1,71 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Build a deterministic punctuation-balanced fine-tuning corpus."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import hashlib
|
||||||
|
import heapq
|
||||||
|
import json
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
CATEGORIES = {
|
||||||
|
"question": lambda text: "?" in text,
|
||||||
|
"exclamation": lambda text: "!" in text,
|
||||||
|
"quotes": lambda text: "«" in text or "»" in text,
|
||||||
|
"ellipsis": lambda text: "…" in text or "..." in text,
|
||||||
|
"neutral": lambda text: not any(mark in text for mark in ("?", "!", "«", "»", "…", "...")),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def build(args: argparse.Namespace) -> None:
|
||||||
|
if args.output.exists() or args.metadata.exists():
|
||||||
|
raise FileExistsError("Refusing to overwrite output or metadata")
|
||||||
|
excluded = set()
|
||||||
|
if args.exclude:
|
||||||
|
excluded = {line.strip() for line in args.exclude.read_text(encoding="utf-8-sig").splitlines() if line.strip()}
|
||||||
|
heaps: dict[str, list[tuple[int, str]]] = {name: [] for name in CATEGORIES}
|
||||||
|
scanned = 0
|
||||||
|
for text in args.source.read_text(encoding="utf-8").splitlines():
|
||||||
|
text = text.strip()
|
||||||
|
if not text or text in excluded:
|
||||||
|
continue
|
||||||
|
scanned += 1
|
||||||
|
for name, predicate in CATEGORIES.items():
|
||||||
|
if not predicate(text):
|
||||||
|
continue
|
||||||
|
score = int.from_bytes(hashlib.blake2b(f"{name}\0{text}".encode("utf-8"), digest_size=8).digest(), "big")
|
||||||
|
heap = heaps[name]
|
||||||
|
if len(heap) < args.per_category:
|
||||||
|
heapq.heappush(heap, (-score, text))
|
||||||
|
elif score < -heap[0][0]:
|
||||||
|
heapq.heapreplace(heap, (-score, text))
|
||||||
|
selected: dict[str, set[str]] = {
|
||||||
|
name: {text for _, text in heap} for name, heap in heaps.items()
|
||||||
|
}
|
||||||
|
combined = sorted(set().union(*selected.values()), key=lambda text: hashlib.sha256(text.encode("utf-8")).digest())
|
||||||
|
args.output.write_text("\n".join(combined) + "\n", encoding="utf-8")
|
||||||
|
metadata = {
|
||||||
|
"source": str(args.source), "exclude": str(args.exclude) if args.exclude else None,
|
||||||
|
"scanned": scanned, "per_category_limit": args.per_category,
|
||||||
|
"category_counts": {name: len(values) for name, values in selected.items()},
|
||||||
|
"unique_sentences": len(combined), "output_bytes": args.output.stat().st_size,
|
||||||
|
}
|
||||||
|
args.metadata.write_text(json.dumps(metadata, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||||
|
print(json.dumps(metadata, ensure_ascii=False, indent=2))
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--source", type=Path, required=True)
|
||||||
|
parser.add_argument("--exclude", type=Path)
|
||||||
|
parser.add_argument("--output", type=Path, required=True)
|
||||||
|
parser.add_argument("--metadata", type=Path, required=True)
|
||||||
|
parser.add_argument("--per-category", type=int, default=4000)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(build(parse_args()) or 0)
|
||||||
@@ -0,0 +1,112 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Compare teacher, FP32 student, and INT8 student embeddings on real text."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import math
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import onnxruntime as ort
|
||||||
|
|
||||||
|
from vosk_booktts_experiment import WordPieceTokenizer, load_corpus, normalize_text
|
||||||
|
|
||||||
|
|
||||||
|
def cosine_rows(left: np.ndarray, right: np.ndarray) -> np.ndarray:
|
||||||
|
numerator = np.sum(left * right, axis=-1)
|
||||||
|
denominator = np.linalg.norm(left, axis=-1) * np.linalg.norm(right, axis=-1)
|
||||||
|
return numerator / np.maximum(denominator, 1e-12)
|
||||||
|
|
||||||
|
|
||||||
|
def compare(args: argparse.Namespace) -> None:
|
||||||
|
tokenizer = WordPieceTokenizer(args.vocab)
|
||||||
|
token_map = np.load(args.token_map)
|
||||||
|
compact_unk = int(token_map[tokenizer.vocabulary["[UNK]"]])
|
||||||
|
sessions = {
|
||||||
|
"teacher": ort.InferenceSession(str(args.teacher), providers=["CPUExecutionProvider"]),
|
||||||
|
"student_fp32": ort.InferenceSession(str(args.student_fp32), providers=["CPUExecutionProvider"]),
|
||||||
|
"student_int8": ort.InferenceSession(str(args.student_int8), providers=["CPUExecutionProvider"]),
|
||||||
|
}
|
||||||
|
timings = {name: 0.0 for name in sessions}
|
||||||
|
teacher_cosines: list[np.ndarray] = []
|
||||||
|
quantized_cosines: list[np.ndarray] = []
|
||||||
|
teacher_squared_error = 0.0
|
||||||
|
quantized_squared_error = 0.0
|
||||||
|
elements = 0
|
||||||
|
tokens = 0
|
||||||
|
unseen_tokens = 0
|
||||||
|
phrases = load_corpus(args.corpus)[: args.limit or None]
|
||||||
|
for text in phrases:
|
||||||
|
ids, _ = tokenizer.encode(normalize_text(text))
|
||||||
|
shape = (1, ids.size)
|
||||||
|
feeds = {
|
||||||
|
"input_ids": ids.reshape(shape),
|
||||||
|
"attention_mask": np.ones(shape, dtype=np.int64),
|
||||||
|
"token_type_ids": np.zeros(shape, dtype=np.int64),
|
||||||
|
}
|
||||||
|
outputs: dict[str, np.ndarray] = {}
|
||||||
|
for name, session in sessions.items():
|
||||||
|
started = time.perf_counter()
|
||||||
|
outputs[name] = session.run(None, feeds)[0].astype(np.float32, copy=False)
|
||||||
|
timings[name] += time.perf_counter() - started
|
||||||
|
teacher = outputs["teacher"]
|
||||||
|
fp32 = outputs["student_fp32"]
|
||||||
|
int8 = outputs["student_int8"]
|
||||||
|
if teacher.shape != fp32.shape or fp32.shape != int8.shape:
|
||||||
|
raise ValueError(f"Output shape mismatch for {text!r}: {teacher.shape}, {fp32.shape}, {int8.shape}")
|
||||||
|
teacher_cosines.append(cosine_rows(teacher, fp32))
|
||||||
|
quantized_cosines.append(cosine_rows(fp32, int8))
|
||||||
|
teacher_squared_error += float(np.sum((teacher - fp32) ** 2))
|
||||||
|
quantized_squared_error += float(np.sum((fp32 - int8) ** 2))
|
||||||
|
elements += teacher.size
|
||||||
|
tokens += ids.size
|
||||||
|
unseen_tokens += int(np.sum((token_map[ids] == compact_unk) & (ids != tokenizer.vocabulary["[UNK]"])))
|
||||||
|
teacher_cosine = np.concatenate(teacher_cosines)
|
||||||
|
quantized_cosine = np.concatenate(quantized_cosines)
|
||||||
|
report = {
|
||||||
|
"phrases": len(phrases),
|
||||||
|
"tokens": tokens,
|
||||||
|
"unseen_tokens": unseen_tokens,
|
||||||
|
"unseen_percent": 100.0 * unseen_tokens / tokens,
|
||||||
|
"teacher_vs_student_fp32": {
|
||||||
|
"mean_token_cosine": float(np.mean(teacher_cosine)),
|
||||||
|
"p05_token_cosine": float(np.percentile(teacher_cosine, 5)),
|
||||||
|
"mse": teacher_squared_error / elements,
|
||||||
|
},
|
||||||
|
"student_fp32_vs_int8": {
|
||||||
|
"mean_token_cosine": float(np.mean(quantized_cosine)),
|
||||||
|
"p05_token_cosine": float(np.percentile(quantized_cosine, 5)),
|
||||||
|
"mse": quantized_squared_error / elements,
|
||||||
|
},
|
||||||
|
"total_inference_seconds": timings,
|
||||||
|
"milliseconds_per_phrase": {name: seconds * 1000.0 / len(phrases) for name, seconds in timings.items()},
|
||||||
|
}
|
||||||
|
if not all(math.isfinite(value) for value in (teacher_squared_error, quantized_squared_error)):
|
||||||
|
raise ValueError("Non-finite comparison result")
|
||||||
|
rendered = json.dumps(report, ensure_ascii=False, indent=2) + "\n"
|
||||||
|
if args.output:
|
||||||
|
if args.output.exists():
|
||||||
|
raise FileExistsError(f"Refusing to overwrite {args.output}")
|
||||||
|
args.output.write_text(rendered, encoding="utf-8")
|
||||||
|
print(rendered, end="")
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--teacher", type=Path, required=True)
|
||||||
|
parser.add_argument("--student-fp32", type=Path, required=True)
|
||||||
|
parser.add_argument("--student-int8", type=Path, required=True)
|
||||||
|
parser.add_argument("--vocab", type=Path, required=True)
|
||||||
|
parser.add_argument("--token-map", type=Path, required=True)
|
||||||
|
parser.add_argument("--corpus", type=Path, required=True)
|
||||||
|
parser.add_argument("--output", type=Path)
|
||||||
|
parser.add_argument("--limit", type=int, default=0)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(compare(parse_args()) or 0)
|
||||||
@@ -0,0 +1,150 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Extract clean Russian sentences from an official Wikisource XML dump."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import bz2
|
||||||
|
import heapq
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import re
|
||||||
|
import sys
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from pathlib import Path
|
||||||
|
from xml.etree import ElementTree
|
||||||
|
|
||||||
|
|
||||||
|
SENTENCE_BOUNDARY = re.compile(r"(?<=[.!?…])\s+")
|
||||||
|
WHITESPACE = re.compile(r"\s+")
|
||||||
|
RUSSIAN_LETTER = re.compile(r"[А-Яа-яЁё]")
|
||||||
|
ANY_LETTER = re.compile(r"[^\W\d_]", re.UNICODE)
|
||||||
|
BAD_MARKUP = ("{{", "}}", "[[", "]]", "http://", "https://", "<math", "</", "{|", "|}")
|
||||||
|
|
||||||
|
|
||||||
|
def local_name(tag: str) -> str:
|
||||||
|
return tag.rsplit("}", 1)[-1]
|
||||||
|
|
||||||
|
|
||||||
|
def descendant_text(element, name: str) -> str:
|
||||||
|
for child in element.iter():
|
||||||
|
if local_name(child.tag) == name:
|
||||||
|
return child.text or ""
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def clean_page_text(source: str) -> str:
|
||||||
|
import mwparserfromhell
|
||||||
|
|
||||||
|
source = re.sub(r"<ref\b[^>]*>.*?</ref>|<ref\b[^>]*/>", " ", source, flags=re.IGNORECASE | re.DOTALL)
|
||||||
|
source = re.sub(r"\{\|.*?\|\}", " ", source, flags=re.DOTALL)
|
||||||
|
plain = mwparserfromhell.parse(source).strip_code(normalize=True, collapse=True)
|
||||||
|
return WHITESPACE.sub(" ", plain).strip()
|
||||||
|
|
||||||
|
|
||||||
|
def acceptable(sentence: str) -> bool:
|
||||||
|
if not 40 <= len(sentence) <= 240 or any(marker in sentence for marker in BAD_MARKUP):
|
||||||
|
return False
|
||||||
|
if sentence[-1:] not in ".!?…":
|
||||||
|
return False
|
||||||
|
words = sentence.split()
|
||||||
|
if not 6 <= len(words) <= 45:
|
||||||
|
return False
|
||||||
|
letters = ANY_LETTER.findall(sentence)
|
||||||
|
if not letters:
|
||||||
|
return False
|
||||||
|
russian = RUSSIAN_LETTER.findall(sentence)
|
||||||
|
return len(russian) / len(letters) >= 0.75
|
||||||
|
|
||||||
|
|
||||||
|
def extract(args: argparse.Namespace) -> None:
|
||||||
|
temporary = args.output.with_suffix(args.output.suffix + ".part")
|
||||||
|
if args.output.exists() or args.metadata.exists() or temporary.exists():
|
||||||
|
raise FileExistsError("Refusing to overwrite output, metadata, or partial output")
|
||||||
|
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
seen: set[bytes] = set()
|
||||||
|
sample_heap: list[tuple[int, bytes, str]] = []
|
||||||
|
pages = 0
|
||||||
|
accepted = 0
|
||||||
|
candidates = 0
|
||||||
|
with bz2.open(args.dump, "rb") as source, temporary.open("w", encoding="utf-8", newline="\n") as output:
|
||||||
|
for _, element in ElementTree.iterparse(source, events=("end",)):
|
||||||
|
if local_name(element.tag) != "page":
|
||||||
|
continue
|
||||||
|
pages += 1
|
||||||
|
namespace = descendant_text(element, "ns")
|
||||||
|
title = descendant_text(element, "title")
|
||||||
|
wiki_text = descendant_text(element, "text")
|
||||||
|
if namespace == "0" and wiki_text and not title.startswith(("Категория:", "Шаблон:", "Справка:")):
|
||||||
|
for sentence in SENTENCE_BOUNDARY.split(clean_page_text(wiki_text)):
|
||||||
|
# Leading dashes carry dialogue/prosody information and must survive extraction.
|
||||||
|
sentence = sentence.strip(" \t\r\n")
|
||||||
|
if not acceptable(sentence):
|
||||||
|
continue
|
||||||
|
digest = hashlib.blake2b(sentence.encode("utf-8"), digest_size=16).digest()
|
||||||
|
candidates += 1
|
||||||
|
if args.sampling == "first":
|
||||||
|
if digest in seen:
|
||||||
|
continue
|
||||||
|
seen.add(digest)
|
||||||
|
output.write(sentence + "\n")
|
||||||
|
accepted += 1
|
||||||
|
if accepted % 10_000 == 0:
|
||||||
|
print(f"sentences={accepted} pages={pages}", flush=True)
|
||||||
|
if accepted >= args.max_sentences:
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
score = int.from_bytes(digest, "big")
|
||||||
|
if digest in seen:
|
||||||
|
continue
|
||||||
|
if len(sample_heap) < args.max_sentences:
|
||||||
|
heapq.heappush(sample_heap, (-score, digest, sentence)); seen.add(digest)
|
||||||
|
elif score < -sample_heap[0][0]:
|
||||||
|
_, removed_digest, _ = heapq.heapreplace(sample_heap, (-score, digest, sentence))
|
||||||
|
seen.remove(removed_digest); seen.add(digest)
|
||||||
|
element.clear()
|
||||||
|
if args.sampling == "first" and accepted >= args.max_sentences:
|
||||||
|
break
|
||||||
|
if pages % 10_000 == 0:
|
||||||
|
print(f"pages={pages} candidates={candidates} retained={len(sample_heap)}", flush=True)
|
||||||
|
if args.sampling == "hash":
|
||||||
|
selected = sorted((-negative_score, sentence) for negative_score, _, sentence in sample_heap)
|
||||||
|
for _, sentence in selected:
|
||||||
|
output.write(sentence + "\n")
|
||||||
|
accepted = len(selected)
|
||||||
|
temporary.replace(args.output)
|
||||||
|
metadata = {
|
||||||
|
"source_url": args.source_url,
|
||||||
|
"dump_file": args.dump.name,
|
||||||
|
"dump_bytes": args.dump.stat().st_size,
|
||||||
|
"pages_scanned": pages,
|
||||||
|
"candidate_sentences": candidates,
|
||||||
|
"sentences": accepted,
|
||||||
|
"output_bytes": args.output.stat().st_size,
|
||||||
|
"created_utc": datetime.now(timezone.utc).isoformat(),
|
||||||
|
"sampling": args.sampling,
|
||||||
|
"filters": {"min_chars": 40, "max_chars": 240, "min_words": 6, "max_words": 45, "min_russian_letter_ratio": 0.75},
|
||||||
|
}
|
||||||
|
args.metadata.write_text(json.dumps(metadata, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||||
|
print(json.dumps(metadata, ensure_ascii=False, indent=2))
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--dump", type=Path, required=True)
|
||||||
|
parser.add_argument("--output", type=Path, required=True)
|
||||||
|
parser.add_argument("--metadata", type=Path, required=True)
|
||||||
|
parser.add_argument("--max-sentences", type=int, default=200_000)
|
||||||
|
parser.add_argument(
|
||||||
|
"--sampling", choices=("first", "hash"), default="hash",
|
||||||
|
help="hash scans the complete dump and keeps a deterministic min-hash sample",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--source-url",
|
||||||
|
default="https://dumps.wikimedia.org/ruwikisource/latest/ruwikisource-latest-pages-articles-multistream.xml.bz2",
|
||||||
|
)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(extract(parse_args()) or 0)
|
||||||
@@ -0,0 +1,198 @@
|
|||||||
|
[CmdletBinding()]
|
||||||
|
param(
|
||||||
|
[int[]] $Milestones = @(100, 300, 500),
|
||||||
|
[int] $PollSeconds = 20
|
||||||
|
)
|
||||||
|
|
||||||
|
$ErrorActionPreference = 'Stop'
|
||||||
|
|
||||||
|
$root = $PSScriptRoot
|
||||||
|
$trainingDir = Join-Path $root 'training'
|
||||||
|
$checkpointRoot = Join-Path $trainingDir 'checkpoints'
|
||||||
|
$configSource = Join-Path $trainingDir 'aletheia_ru.onnx.json'
|
||||||
|
$python = Join-Path $root '.venv\Scripts\python.exe'
|
||||||
|
$outputRoot = Join-Path $trainingDir 'milestones'
|
||||||
|
$monitorLog = Join-Path $outputRoot 'monitor.log'
|
||||||
|
New-Item -ItemType Directory -Force -Path $outputRoot | Out-Null
|
||||||
|
|
||||||
|
function Write-MonitorLog([string] $Message) {
|
||||||
|
$line = '{0} {1}' -f (Get-Date).ToUniversalTime().ToString('o'), $Message
|
||||||
|
Add-Content -LiteralPath $monitorLog -Value $line -Encoding utf8
|
||||||
|
}
|
||||||
|
|
||||||
|
function Invoke-LoggedProcess(
|
||||||
|
[string] $FilePath,
|
||||||
|
[string[]] $Arguments,
|
||||||
|
[string] $StdoutPath,
|
||||||
|
[string] $StderrPath
|
||||||
|
) {
|
||||||
|
$process = Start-Process -FilePath $FilePath `
|
||||||
|
-ArgumentList $Arguments `
|
||||||
|
-RedirectStandardOutput $StdoutPath `
|
||||||
|
-RedirectStandardError $StderrPath `
|
||||||
|
-WindowStyle Hidden `
|
||||||
|
-Wait `
|
||||||
|
-PassThru
|
||||||
|
if ($process.ExitCode -ne 0) {
|
||||||
|
throw "Process failed with exit code $($process.ExitCode): $FilePath"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function Get-LatestCheckpoint {
|
||||||
|
Get-ChildItem -LiteralPath $checkpointRoot -Filter '*.ckpt' -File -Recurse -ErrorAction SilentlyContinue |
|
||||||
|
Where-Object { $_.Name -match '^epoch=(\d+)-step=(\d+)\.ckpt$' } |
|
||||||
|
Sort-Object LastWriteTimeUtc -Descending |
|
||||||
|
Select-Object -First 1
|
||||||
|
}
|
||||||
|
|
||||||
|
function Show-TrainingStatus {
|
||||||
|
$trainingLog = Join-Path $trainingDir 'scheduled-training.stdout.log'
|
||||||
|
$trainingTask = Get-ScheduledTask -TaskName 'AletheiaTTS-Training' -ErrorAction SilentlyContinue
|
||||||
|
$progressLine = ''
|
||||||
|
if (Test-Path -LiteralPath $trainingLog -PathType Leaf) {
|
||||||
|
$progressLine = Get-Content -LiteralPath $trainingLog -Tail 30 |
|
||||||
|
Where-Object { $_ -match '^Epoch\s+\d+:' } |
|
||||||
|
Select-Object -Last 1
|
||||||
|
}
|
||||||
|
$checkpoint = Get-LatestCheckpoint
|
||||||
|
$gpu = (& nvidia-smi --query-gpu=name,utilization.gpu,memory.used,memory.total,temperature.gpu,power.draw --format=csv,noheader,nounits 2>$null) -join [Environment]::NewLine
|
||||||
|
|
||||||
|
try {
|
||||||
|
$Host.UI.RawUI.WindowTitle = 'Aletheia TTS - training monitor'
|
||||||
|
Clear-Host
|
||||||
|
} catch {
|
||||||
|
# The monitor also works when no interactive console is attached.
|
||||||
|
}
|
||||||
|
Write-Host 'Aletheia Russian TTS training'
|
||||||
|
Write-Host ('Updated: {0}' -f (Get-Date).ToString('yyyy-MM-dd HH:mm:ss'))
|
||||||
|
Write-Host ('Training task: {0}' -f $(if ($trainingTask) { $trainingTask.State } else { 'Not found' }))
|
||||||
|
Write-Host ('Progress: {0}' -f $(if ($progressLine) { $progressLine.Trim() } else { 'Waiting for log data' }))
|
||||||
|
Write-Host ('GPU: {0}' -f $(if ($gpu) { $gpu.Trim() } else { 'nvidia-smi unavailable' }))
|
||||||
|
if ($checkpoint) {
|
||||||
|
Write-Host ('Latest checkpoint: {0} ({1:N0} bytes)' -f $checkpoint.Name, $checkpoint.Length)
|
||||||
|
} else {
|
||||||
|
Write-Host 'Latest checkpoint: waiting'
|
||||||
|
}
|
||||||
|
Write-Host ''
|
||||||
|
Write-Host 'Evaluation milestones:'
|
||||||
|
foreach ($milestone in $Milestones) {
|
||||||
|
$completePath = Join-Path $outputRoot ('epoch-{0:d4}\complete.json' -f $milestone)
|
||||||
|
$state = if (Test-Path -LiteralPath $completePath -PathType Leaf) { 'READY' } else { 'waiting' }
|
||||||
|
Write-Host (' {0,4} epochs: {1}' -f $milestone, $state)
|
||||||
|
}
|
||||||
|
Write-Host ''
|
||||||
|
Write-Host 'This window may be minimized. Closing it stops milestone exports only.'
|
||||||
|
}
|
||||||
|
|
||||||
|
function Save-Milestone([int] $Milestone, [IO.FileInfo] $SourceCheckpoint) {
|
||||||
|
$milestoneName = 'epoch-{0:d4}' -f $Milestone
|
||||||
|
$milestoneDir = Join-Path $outputRoot $milestoneName
|
||||||
|
$completePath = Join-Path $milestoneDir 'complete.json'
|
||||||
|
if (Test-Path -LiteralPath $completePath -PathType Leaf) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
New-Item -ItemType Directory -Force -Path $milestoneDir | Out-Null
|
||||||
|
$checkpointPath = Join-Path $milestoneDir "aletheia_ru_$milestoneName.ckpt"
|
||||||
|
if (-not (Test-Path -LiteralPath $checkpointPath -PathType Leaf)) {
|
||||||
|
$firstLength = $SourceCheckpoint.Length
|
||||||
|
Start-Sleep -Seconds 10
|
||||||
|
$refreshed = Get-Item -LiteralPath $SourceCheckpoint.FullName
|
||||||
|
if ($firstLength -ne $refreshed.Length -or $refreshed.Length -le 0) {
|
||||||
|
throw "Checkpoint is not stable yet: $($SourceCheckpoint.FullName)"
|
||||||
|
}
|
||||||
|
$partialPath = "$checkpointPath.partial"
|
||||||
|
Copy-Item -LiteralPath $refreshed.FullName -Destination $partialPath -Force
|
||||||
|
Move-Item -LiteralPath $partialPath -Destination $checkpointPath -Force
|
||||||
|
}
|
||||||
|
|
||||||
|
$modelPath = Join-Path $milestoneDir "aletheia_ru_$milestoneName.onnx"
|
||||||
|
$configPath = "$modelPath.json"
|
||||||
|
if (-not (Test-Path -LiteralPath $modelPath -PathType Leaf)) {
|
||||||
|
Invoke-LoggedProcess $python @(
|
||||||
|
'-m', 'piper.train.export_onnx',
|
||||||
|
'--checkpoint', $checkpointPath,
|
||||||
|
'--output-file', $modelPath
|
||||||
|
) (Join-Path $milestoneDir 'export.stdout.log') (Join-Path $milestoneDir 'export.stderr.log')
|
||||||
|
}
|
||||||
|
Copy-Item -LiteralPath $configSource -Destination $configPath -Force
|
||||||
|
|
||||||
|
$inputPath = Join-Path $milestoneDir 'test_sentences.txt'
|
||||||
|
$testText = @(
|
||||||
|
'0JIg0YLQuNGI0LjQvdC1INCy0LXRh9C10YDQvdC10Lkg0LHQuNCx0LvQuNC+0YLQtdC60Lgg0YjQtdC70LXRgdGC0LXQu9C4INGB0YLRgNCw0L3QuNGG0YssINC4INC60LDQttC00LDRjyDQvdC+0LLQsNGPINCz0LvQsNCy0LAg0L7RgtC60YDRi9Cy0LDQu9CwINGD0LTQuNCy0LjRgtC10LvRjNC90YvQuSDQvNC40YAu',
|
||||||
|
'0JrQvtCz0LTQsCDRh9C10LvQvtCy0LXQuiDQtNC10LvQsNC10YIg0LLRi9Cx0L7RgCwg0L7QvSDQvdC1INCy0YHQtdCz0LTQsCDQt9Cw0YDQsNC90LXQtSDQt9C90LDQtdGCLCDQuiDQutCw0LrQuNC8INC/0L7RgdC70LXQtNGB0YLQstC40Y/QvCDQv9GA0LjQstC10LTRkdGCINC10LPQviDRgNC10YjQtdC90LjQtS4=',
|
||||||
|
'0JfQsCDQvtC60L3QvtC8INC80LXQtNC70LXQvdC90L4g0L3QsNGH0LjQvdCw0LvRgdGPINC00L7QttC00YwsINC90L4g0L/Rg9GC0LXRiNC10YHRgtCy0LjQtSDQs9C10YDQvtC10LIg0YLQvtC70YzQutC+INC90LDQsdC40YDQsNC70L4g0YHQuNC70YMu'
|
||||||
|
) | ForEach-Object { [Text.Encoding]::UTF8.GetString([Convert]::FromBase64String($_)) }
|
||||||
|
$testText = $testText -join [Environment]::NewLine
|
||||||
|
[IO.File]::WriteAllText($inputPath, $testText, [Text.UTF8Encoding]::new($false))
|
||||||
|
|
||||||
|
$samples = @(
|
||||||
|
@{ Name = 'neutral'; Length = '1.0'; Noise = '0.667'; Width = '0.8'; Silence = '0.0' },
|
||||||
|
@{ Name = 'book'; Length = '1.5'; Noise = '0.75'; Width = '0.95'; Silence = '0.25' },
|
||||||
|
@{ Name = 'fast'; Length = '0.82'; Noise = '0.70'; Width = '0.85'; Silence = '0.0' }
|
||||||
|
)
|
||||||
|
foreach ($sample in $samples) {
|
||||||
|
$wavPath = Join-Path $milestoneDir ("sample_{0}.wav" -f $sample.Name)
|
||||||
|
if (Test-Path -LiteralPath $wavPath -PathType Leaf) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
Invoke-LoggedProcess $python @(
|
||||||
|
'-m', 'piper',
|
||||||
|
'--model', $modelPath,
|
||||||
|
'--config', $configPath,
|
||||||
|
'--input-file', $inputPath,
|
||||||
|
'--output-file', $wavPath,
|
||||||
|
'--length-scale', $sample.Length,
|
||||||
|
'--noise-scale', $sample.Noise,
|
||||||
|
'--noise-w-scale', $sample.Width,
|
||||||
|
'--sentence-silence', $sample.Silence
|
||||||
|
) (Join-Path $milestoneDir "$($sample.Name).stdout.log") (Join-Path $milestoneDir "$($sample.Name).stderr.log")
|
||||||
|
}
|
||||||
|
|
||||||
|
$sourceMatch = [regex]::Match($SourceCheckpoint.Name, '^epoch=(\d+)-step=(\d+)\.ckpt$')
|
||||||
|
$artifacts = Get-ChildItem -LiteralPath $milestoneDir -File |
|
||||||
|
Where-Object { $_.Extension -in @('.ckpt', '.onnx', '.json', '.wav') } |
|
||||||
|
ForEach-Object {
|
||||||
|
[ordered]@{
|
||||||
|
name = $_.Name
|
||||||
|
bytes = $_.Length
|
||||||
|
sha256 = (Get-FileHash -LiteralPath $_.FullName -Algorithm SHA256).Hash.ToLowerInvariant()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
[ordered]@{
|
||||||
|
requested_completed_epochs = $Milestone
|
||||||
|
source_epoch_zero_based = [int]$sourceMatch.Groups[1].Value
|
||||||
|
source_step = [int]$sourceMatch.Groups[2].Value
|
||||||
|
created_utc = (Get-Date).ToUniversalTime().ToString('o')
|
||||||
|
artifacts = @($artifacts)
|
||||||
|
} | ConvertTo-Json -Depth 5 | Set-Content -LiteralPath $completePath -Encoding utf8
|
||||||
|
Write-MonitorLog "Completed milestone $Milestone from $($SourceCheckpoint.Name)"
|
||||||
|
}
|
||||||
|
|
||||||
|
Write-MonitorLog "Monitor started for milestones: $($Milestones -join ', ')"
|
||||||
|
while ($true) {
|
||||||
|
Show-TrainingStatus
|
||||||
|
$pending = @($Milestones | Where-Object {
|
||||||
|
-not (Test-Path -LiteralPath (Join-Path $outputRoot ('epoch-{0:d4}\complete.json' -f $_)) -PathType Leaf)
|
||||||
|
})
|
||||||
|
if ($pending.Count -eq 0) {
|
||||||
|
Write-MonitorLog 'All milestones completed'
|
||||||
|
exit 0
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
$checkpoint = Get-LatestCheckpoint
|
||||||
|
if ($checkpoint -and $checkpoint.Name -match '^epoch=(\d+)-step=(\d+)\.ckpt$') {
|
||||||
|
$completedEpochs = [int]$Matches[1] + 1
|
||||||
|
foreach ($milestone in $pending) {
|
||||||
|
if ($completedEpochs -ge $milestone) {
|
||||||
|
Save-Milestone $milestone $checkpoint
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} catch {
|
||||||
|
Write-MonitorLog ("Retryable error: " + ($_ | Out-String).Trim())
|
||||||
|
}
|
||||||
|
|
||||||
|
Start-Sleep -Seconds $PollSeconds
|
||||||
|
}
|
||||||
@@ -0,0 +1,54 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Add the original-to-compact token lookup to an already exported student ONNX."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import onnx
|
||||||
|
from onnx import helper, numpy_helper
|
||||||
|
|
||||||
|
|
||||||
|
def patch(args: argparse.Namespace) -> None:
|
||||||
|
if args.output.exists():
|
||||||
|
raise FileExistsError(f"Refusing to overwrite {args.output}")
|
||||||
|
model = onnx.load(str(args.model))
|
||||||
|
if not any(value.name == "input_ids" for value in model.graph.input):
|
||||||
|
raise ValueError("Model has no input_ids graph input")
|
||||||
|
if any(item.name == "original_to_compact_token_id" for item in model.graph.initializer):
|
||||||
|
raise ValueError("Model already contains original_to_compact_token_id")
|
||||||
|
consumers = 0
|
||||||
|
for node in model.graph.node:
|
||||||
|
for index, name in enumerate(node.input):
|
||||||
|
if name == "input_ids":
|
||||||
|
node.input[index] = "mapped_input_ids"
|
||||||
|
consumers += 1
|
||||||
|
if consumers == 0:
|
||||||
|
raise ValueError("No input_ids consumers found")
|
||||||
|
token_map = np.load(args.token_map).astype(np.int64, copy=False)
|
||||||
|
model.graph.initializer.append(numpy_helper.from_array(token_map, "original_to_compact_token_id"))
|
||||||
|
model.graph.node.insert(
|
||||||
|
0,
|
||||||
|
helper.make_node(
|
||||||
|
"Gather", ["original_to_compact_token_id", "input_ids"], ["mapped_input_ids"],
|
||||||
|
axis=0, name="MapOriginalTokenIds",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
onnx.checker.check_model(model)
|
||||||
|
onnx.save(model, str(args.output))
|
||||||
|
print(f"token_map_entries={token_map.size} consumers={consumers} output_bytes={args.output.stat().st_size}")
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--model", type=Path, required=True)
|
||||||
|
parser.add_argument("--token-map", type=Path, required=True)
|
||||||
|
parser.add_argument("--output", type=Path, required=True)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(patch(parse_args()) or 0)
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
diff --git a/src/piper/train/export_onnx.py b/src/piper/train/export_onnx.py
|
||||||
|
index fe9be60..abfeeb0 100644
|
||||||
|
--- a/src/piper/train/export_onnx.py
|
||||||
|
+++ b/src/piper/train/export_onnx.py
|
||||||
|
@@ -103,5 +103,6 @@ def main() -> None:
|
||||||
|
"input_lengths": {0: "batch_size"},
|
||||||
|
"output": {0: "batch_size", 2: "time"},
|
||||||
|
},
|
||||||
|
+ dynamo=False,
|
||||||
|
)
|
||||||
|
_LOGGER.info("Exported model to %s", output_path)
|
||||||
@@ -0,0 +1,78 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Prepare a Piper ONNX export for sherpa-onnx mobile inference."""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import shutil
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import onnx
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--model", type=Path, required=True)
|
||||||
|
parser.add_argument("--config", type=Path, required=True)
|
||||||
|
parser.add_argument("--output-dir", type=Path, required=True)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
config = json.loads(args.config.read_text(encoding="utf-8"))
|
||||||
|
output_dir: Path = args.output_dir
|
||||||
|
output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
output_model = output_dir / "aletheia_ru.onnx"
|
||||||
|
output_config = output_dir / "aletheia_ru.onnx.json"
|
||||||
|
output_tokens = output_dir / "tokens.txt"
|
||||||
|
|
||||||
|
shutil.copy2(args.model, output_model)
|
||||||
|
shutil.copy2(args.config, output_config)
|
||||||
|
|
||||||
|
id_map = config["phoneme_id_map"]
|
||||||
|
# sherpa-onnx's Piper lexicon maps one Unicode code point to one model ID.
|
||||||
|
# Newer Piper configs may also contain English diphthong aliases such as
|
||||||
|
# "aɪ". eSpeak emits their component code points, and sherpa rejects the
|
||||||
|
# multi-code-point aliases, so only the single-code-point table is mobile-safe.
|
||||||
|
token_rows = sorted(
|
||||||
|
((ids[0], symbol) for symbol, ids in id_map.items() if len(symbol) == 1),
|
||||||
|
key=lambda row: row[0],
|
||||||
|
)
|
||||||
|
actual_ids = [row[0] for row in token_rows]
|
||||||
|
if len(actual_ids) != len(set(actual_ids)):
|
||||||
|
raise RuntimeError("phoneme_id_map contains duplicate token identifiers")
|
||||||
|
if not actual_ids or min(actual_ids) < 0 or max(actual_ids) >= config["num_symbols"]:
|
||||||
|
raise RuntimeError("phoneme_id_map contains a token outside the model symbol range")
|
||||||
|
output_tokens.write_text(
|
||||||
|
"".join(f"{symbol} {token_id}\n" for token_id, symbol in token_rows),
|
||||||
|
encoding="utf-8",
|
||||||
|
newline="\n",
|
||||||
|
)
|
||||||
|
|
||||||
|
model = onnx.load(str(output_model))
|
||||||
|
metadata = {
|
||||||
|
"model_type": "vits",
|
||||||
|
"comment": "piper",
|
||||||
|
"language": "Russian",
|
||||||
|
"voice": config["espeak"]["voice"],
|
||||||
|
"has_espeak": "1",
|
||||||
|
"n_speakers": str(config["num_speakers"]),
|
||||||
|
"sample_rate": str(config["audio"]["sample_rate"]),
|
||||||
|
}
|
||||||
|
existing = {item.key: item for item in model.metadata_props}
|
||||||
|
for key, value in metadata.items():
|
||||||
|
if key in existing:
|
||||||
|
existing[key].value = value
|
||||||
|
else:
|
||||||
|
item = model.metadata_props.add()
|
||||||
|
item.key = key
|
||||||
|
item.value = value
|
||||||
|
onnx.save(model, str(output_model))
|
||||||
|
|
||||||
|
print(json.dumps({
|
||||||
|
"model": str(output_model),
|
||||||
|
"model_bytes": output_model.stat().st_size,
|
||||||
|
"tokens": len(token_rows),
|
||||||
|
"metadata": metadata,
|
||||||
|
}, ensure_ascii=False))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
param([string]$Root = 'C:\Users\seven\AletheiaBookTTS')
|
||||||
|
|
||||||
|
$ErrorActionPreference = 'Stop'
|
||||||
|
$env:PYTHONUTF8 = '1'
|
||||||
|
$env:PYTHONIOENCODING = 'utf-8'
|
||||||
|
$python = Join-Path $Root '.venv\Scripts\python.exe'
|
||||||
|
$work = Join-Path $Root 'work'
|
||||||
|
$candidate = Join-Path $work 'student-100k-v1'
|
||||||
|
& $python (Join-Path $work 'vosk_booktts_experiment.py') compare-feeds `
|
||||||
|
--model (Join-Path $Root 'teacher\model.onnx') `
|
||||||
|
--baseline-calibration (Join-Path $candidate 'feeds-baseline-100-v2') `
|
||||||
|
--candidate-calibration (Join-Path $candidate 'feeds-student-int8-100-v2') `
|
||||||
|
--output (Join-Path $candidate 'wav-ab-100-v1') `
|
||||||
|
--resume
|
||||||
|
exit $LASTEXITCODE
|
||||||
@@ -0,0 +1,89 @@
|
|||||||
|
param(
|
||||||
|
[string]$Root = 'C:\Users\seven\AletheiaBookTTS',
|
||||||
|
[int]$MaxSentences = 100000,
|
||||||
|
[int]$TeacherBatchSize = 16,
|
||||||
|
[ValidateSet('first', 'hash')]
|
||||||
|
[string]$Sampling = 'hash'
|
||||||
|
)
|
||||||
|
|
||||||
|
$ErrorActionPreference = 'Stop'
|
||||||
|
$ProgressPreference = 'SilentlyContinue'
|
||||||
|
$env:PYTHONUTF8 = '1'
|
||||||
|
$env:PYTHONIOENCODING = 'utf-8'
|
||||||
|
|
||||||
|
$displayName = 'AletheiaBookTTS-RuWikisource'
|
||||||
|
$corpusDir = Join-Path $Root 'corpus'
|
||||||
|
$workDir = Join-Path $Root 'work'
|
||||||
|
$teacherDir = Join-Path $Root 'teacher'
|
||||||
|
$python = Join-Path $Root '.venv\Scripts\python.exe'
|
||||||
|
$dump = Join-Path $corpusDir 'ruwikisource-latest-pages-articles-multistream.xml.bz2'
|
||||||
|
$sentences = Join-Path $corpusDir "ruwikisource-booktts-$MaxSentences-$Sampling.txt"
|
||||||
|
$metadata = Join-Path $corpusDir "ruwikisource-booktts-$MaxSentences-$Sampling.json"
|
||||||
|
$dataset = Join-Path $workDir "student-dataset-$MaxSentences-$Sampling"
|
||||||
|
$datasetPartial = "$dataset.partial"
|
||||||
|
$expectedBytes = 2104816879L
|
||||||
|
$checksumUrl = 'https://dumps.wikimedia.org/ruwikisource/latest/ruwikisource-latest-sha1sums.txt'
|
||||||
|
|
||||||
|
Import-Module BitsTransfer
|
||||||
|
while (-not (Test-Path -LiteralPath $dump)) {
|
||||||
|
$job = Get-BitsTransfer -ErrorAction SilentlyContinue |
|
||||||
|
Where-Object DisplayName -eq $displayName |
|
||||||
|
Select-Object -First 1
|
||||||
|
if (-not $job) {
|
||||||
|
throw "BITS job $displayName is absent and dump is not downloaded"
|
||||||
|
}
|
||||||
|
if ($job.JobState -eq 'Transferred') {
|
||||||
|
Complete-BitsTransfer -BitsJob $job
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if ($job.JobState -eq 'Error') {
|
||||||
|
throw "BITS download failed: $($job.ErrorDescription)"
|
||||||
|
}
|
||||||
|
$percent = if ($job.BytesTotal -gt 0 -and $job.BytesTotal -lt [uint64]::MaxValue) {
|
||||||
|
[math]::Round(100 * $job.BytesTransferred / $job.BytesTotal, 2)
|
||||||
|
} else { 0 }
|
||||||
|
Write-Output "download state=$($job.JobState) bytes=$($job.BytesTransferred)/$($job.BytesTotal) percent=$percent"
|
||||||
|
Start-Sleep -Seconds 20
|
||||||
|
}
|
||||||
|
|
||||||
|
$dumpItem = Get-Item -LiteralPath $dump
|
||||||
|
if ($dumpItem.Length -ne $expectedBytes) {
|
||||||
|
throw "Unexpected dump size: $($dumpItem.Length), expected $expectedBytes"
|
||||||
|
}
|
||||||
|
$checksumText = (Invoke-WebRequest -UseBasicParsing -Uri $checksumUrl).Content
|
||||||
|
$checksumLine = ($checksumText -split "`n" | Where-Object { $_ -match 'pages-articles-multistream\.xml\.bz2\s*$' } | Select-Object -First 1).Trim()
|
||||||
|
if (-not $checksumLine) {
|
||||||
|
throw 'Cannot find multistream dump in official SHA-1 list'
|
||||||
|
}
|
||||||
|
$expectedSha1 = ($checksumLine -split '\s+')[0].ToUpperInvariant()
|
||||||
|
$actualSha1 = (Get-FileHash -LiteralPath $dump -Algorithm SHA1).Hash
|
||||||
|
if ($actualSha1 -ne $expectedSha1) {
|
||||||
|
throw "SHA-1 mismatch: actual=$actualSha1 expected=$expectedSha1"
|
||||||
|
}
|
||||||
|
Write-Output "dump verified bytes=$($dumpItem.Length) sha1=$actualSha1"
|
||||||
|
|
||||||
|
if (-not (Test-Path -LiteralPath $sentences)) {
|
||||||
|
& $python (Join-Path $workDir 'extract_wikisource_corpus.py') `
|
||||||
|
--dump $dump `
|
||||||
|
--output $sentences `
|
||||||
|
--metadata $metadata `
|
||||||
|
--max-sentences $MaxSentences `
|
||||||
|
--sampling $Sampling
|
||||||
|
if ($LASTEXITCODE -ne 0) { throw "Corpus extraction failed with exit code $LASTEXITCODE" }
|
||||||
|
}
|
||||||
|
|
||||||
|
if (-not (Test-Path -LiteralPath $dataset)) {
|
||||||
|
if (Test-Path -LiteralPath $datasetPartial) {
|
||||||
|
throw "Partial dataset already exists: $datasetPartial"
|
||||||
|
}
|
||||||
|
& $python (Join-Path $workDir 'booktts_student.py') dataset `
|
||||||
|
--assets $teacherDir `
|
||||||
|
--corpus $sentences `
|
||||||
|
--output $datasetPartial `
|
||||||
|
--teacher-batch-size $TeacherBatchSize `
|
||||||
|
--teacher-dtype float16
|
||||||
|
if ($LASTEXITCODE -ne 0) { throw "Teacher dataset failed with exit code $LASTEXITCODE" }
|
||||||
|
Move-Item -LiteralPath $datasetPartial -Destination $dataset
|
||||||
|
}
|
||||||
|
|
||||||
|
Write-Output "pipeline complete corpus=$sentences dataset=$dataset"
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
param([string]$Root = 'C:\Users\seven\AletheiaBookTTS')
|
||||||
|
|
||||||
|
$ErrorActionPreference = 'Stop'
|
||||||
|
$env:PYTHONUTF8 = '1'
|
||||||
|
$env:PYTHONIOENCODING = 'utf-8'
|
||||||
|
$python = Join-Path $Root '.venv\Scripts\python.exe'
|
||||||
|
$work = Join-Path $Root 'work'
|
||||||
|
$output = Join-Path $work 'student-dataset-prosody-v1'
|
||||||
|
& $python (Join-Path $work 'booktts_student.py') dataset `
|
||||||
|
--assets (Join-Path $Root 'teacher') `
|
||||||
|
--corpus (Join-Path $Root 'corpus\booktts-prosody-balanced-v1.txt') `
|
||||||
|
--output $output `
|
||||||
|
--teacher-batch-size 16 `
|
||||||
|
--teacher-dtype float16 `
|
||||||
|
--reuse-token-map (Join-Path $work 'student-dataset-100000')
|
||||||
|
exit $LASTEXITCODE
|
||||||
@@ -0,0 +1,17 @@
|
|||||||
|
param([string]$Root = 'C:\Users\seven\AletheiaBookTTS')
|
||||||
|
|
||||||
|
$ErrorActionPreference = 'Stop'
|
||||||
|
$env:PYTHONUTF8 = '1'
|
||||||
|
$env:PYTHONIOENCODING = 'utf-8'
|
||||||
|
$python = Join-Path $Root '.venv\Scripts\python.exe'
|
||||||
|
$work = Join-Path $Root 'work'
|
||||||
|
& $python (Join-Path $work 'booktts_student.py') train `
|
||||||
|
--dataset (Join-Path $work 'student-dataset-prosody-v1') `
|
||||||
|
--initial-model (Join-Path $work 'student-100k-v1\student.pt') `
|
||||||
|
--output (Join-Path $work 'student-100k-prosody-v1') `
|
||||||
|
--epochs 8 `
|
||||||
|
--batch-size 32 `
|
||||||
|
--learning-rate 0.00005 `
|
||||||
|
--validation-fraction 0.1 `
|
||||||
|
--early-stopping-patience 3
|
||||||
|
exit $LASTEXITCODE
|
||||||
@@ -0,0 +1,27 @@
|
|||||||
|
param(
|
||||||
|
[string]$Root = 'C:\Users\seven\AletheiaBookTTS',
|
||||||
|
[string]$DatasetName = 'student-dataset-100000',
|
||||||
|
[string]$OutputName = 'student-100k-v1',
|
||||||
|
[int]$Epochs = 20,
|
||||||
|
[int]$BatchSize = 32,
|
||||||
|
[double]$ValidationFraction = 0.05,
|
||||||
|
[int]$EarlyStoppingPatience = 4,
|
||||||
|
[switch]$Resume
|
||||||
|
)
|
||||||
|
|
||||||
|
$ErrorActionPreference = 'Stop'
|
||||||
|
$env:PYTHONUTF8 = '1'
|
||||||
|
$env:PYTHONIOENCODING = 'utf-8'
|
||||||
|
$python = Join-Path $Root '.venv\Scripts\python.exe'
|
||||||
|
$trainer = Join-Path $Root 'work\booktts_student.py'
|
||||||
|
$dataset = Join-Path $Root "work\$DatasetName"
|
||||||
|
$output = Join-Path $Root "work\$OutputName"
|
||||||
|
$arguments = @(
|
||||||
|
$trainer, 'train', '--dataset', $dataset, '--output', $output,
|
||||||
|
'--epochs', $Epochs, '--batch-size', $BatchSize,
|
||||||
|
'--validation-fraction', $ValidationFraction,
|
||||||
|
'--early-stopping-patience', $EarlyStoppingPatience
|
||||||
|
)
|
||||||
|
if ($Resume) { $arguments += '--resume' }
|
||||||
|
& $python @arguments
|
||||||
|
exit $LASTEXITCODE
|
||||||
@@ -15,6 +15,13 @@ if (Test-Path -LiteralPath $exitPath) {
|
|||||||
|
|
||||||
$exitCode = 1
|
$exitCode = 1
|
||||||
try {
|
try {
|
||||||
|
$latestCheckpoint = Get-ChildItem -LiteralPath (Join-Path $runDir 'checkpoints') `
|
||||||
|
-Filter '*.ckpt' `
|
||||||
|
-File `
|
||||||
|
-Recurse `
|
||||||
|
-ErrorAction SilentlyContinue |
|
||||||
|
Sort-Object LastWriteTimeUtc -Descending |
|
||||||
|
Select-Object -First 1
|
||||||
$arguments = @(
|
$arguments = @(
|
||||||
'-NoProfile', '-NonInteractive', '-ExecutionPolicy', 'Bypass',
|
'-NoProfile', '-NonInteractive', '-ExecutionPolicy', 'Bypass',
|
||||||
'-File', (Join-Path $root 'train_piper_windows.ps1'),
|
'-File', (Join-Path $root 'train_piper_windows.ps1'),
|
||||||
@@ -24,6 +31,9 @@ try {
|
|||||||
'-PythonExe', (Join-Path $root '.venv\Scripts\python.exe'),
|
'-PythonExe', (Join-Path $root '.venv\Scripts\python.exe'),
|
||||||
'-BatchSize', '16', '-NumWorkers', '0', '-MaxEpochs', '2000'
|
'-BatchSize', '16', '-NumWorkers', '0', '-MaxEpochs', '2000'
|
||||||
)
|
)
|
||||||
|
if ($latestCheckpoint) {
|
||||||
|
$arguments += @('-CheckpointPath', $latestCheckpoint.FullName)
|
||||||
|
}
|
||||||
$process = Start-Process -FilePath 'powershell.exe' `
|
$process = Start-Process -FilePath 'powershell.exe' `
|
||||||
-ArgumentList $arguments `
|
-ArgumentList $arguments `
|
||||||
-RedirectStandardOutput $stdoutPath `
|
-RedirectStandardOutput $stdoutPath `
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ param(
|
|||||||
[Parameter(Mandatory = $true)]
|
[Parameter(Mandatory = $true)]
|
||||||
[string] $PythonExe,
|
[string] $PythonExe,
|
||||||
[string] $CacheDir,
|
[string] $CacheDir,
|
||||||
|
[string] $CheckpointPath,
|
||||||
[int] $BatchSize = 4,
|
[int] $BatchSize = 4,
|
||||||
[int] $NumWorkers = 2,
|
[int] $NumWorkers = 2,
|
||||||
[int] $MaxEpochs = 2000,
|
[int] $MaxEpochs = 2000,
|
||||||
@@ -61,6 +62,13 @@ $fitArgs = @(
|
|||||||
if ($SmokeTest) {
|
if ($SmokeTest) {
|
||||||
$fitArgs += @('--trainer.fast_dev_run', 'true', '--trainer.num_sanity_val_steps', '0')
|
$fitArgs += @('--trainer.fast_dev_run', 'true', '--trainer.num_sanity_val_steps', '0')
|
||||||
}
|
}
|
||||||
|
if ($CheckpointPath) {
|
||||||
|
$checkpointToResume = [IO.Path]::GetFullPath($CheckpointPath)
|
||||||
|
if (-not (Test-Path -LiteralPath $checkpointToResume -PathType Leaf)) {
|
||||||
|
throw "Resume checkpoint not found: $checkpointToResume"
|
||||||
|
}
|
||||||
|
$fitArgs += @('--ckpt_path', $checkpointToResume)
|
||||||
|
}
|
||||||
|
|
||||||
& $python @fitArgs
|
& $python @fitArgs
|
||||||
if ($LASTEXITCODE -ne 0) {
|
if ($LASTEXITCODE -ne 0) {
|
||||||
|
|||||||
@@ -0,0 +1,568 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Prepare real Vosk feeds, quantize acoustic ONNX, and render A/B WAVs.
|
||||||
|
|
||||||
|
The frontend mirrors VoskSpeechEngine.kt: number normalization, WordPiece
|
||||||
|
positions, dictionary lookup and all five phone feature channels.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import bisect
|
||||||
|
import json
|
||||||
|
import math
|
||||||
|
import re
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
import wave
|
||||||
|
from collections import OrderedDict
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Iterator
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
SAMPLE_RATE = 22_050
|
||||||
|
SPEAKER_ID = 4
|
||||||
|
BERT_DIMENSIONS = 768
|
||||||
|
MAX_WORD_PIECE_CHARS = 100
|
||||||
|
CACHE_SIZE = 4_096
|
||||||
|
MULTISTREAM_PUNCTUATION = set("!(),-.:;?")
|
||||||
|
BERT_EXCLUDED_PUNCTUATION = {"-", ",", ".", "?", "!", ";", ":", '"'}
|
||||||
|
PHONEMES = [
|
||||||
|
"_", "^", "$", " ", "!", "'", "(", ")", ",", "-", ".", "...", ":", ";", "?",
|
||||||
|
"a0", "a1", "b", "bj", "c", "ch", "d", "dj", "e0", "e1", "f", "fj", "g", "gj",
|
||||||
|
"h", "hj", "i0", "i1", "j", "k", "kj", "l", "lj", "m", "mj", "n", "nj", "o0",
|
||||||
|
"o1", "p", "pj", "r", "rj", "s", "sch", "sh", "sj", "t", "tj", "u0", "u1", "v",
|
||||||
|
"vj", "y0", "y1", "z", "zh", "zj",
|
||||||
|
]
|
||||||
|
PHONEME_IDS = {phone: index for index, phone in enumerate(PHONEMES)}
|
||||||
|
|
||||||
|
DIGITS = ["ноль", "один", "два", "три", "четыре", "пять", "шесть", "семь", "восемь", "девять"]
|
||||||
|
TEENS = [
|
||||||
|
"десять", "одиннадцать", "двенадцать", "тринадцать", "четырнадцать",
|
||||||
|
"пятнадцать", "шестнадцать", "семнадцать", "восемнадцать", "девятнадцать",
|
||||||
|
]
|
||||||
|
TENS = ["", "десять", "двадцать", "тридцать", "сорок", "пятьдесят", "шестьдесят", "семьдесят", "восемьдесят", "девяносто"]
|
||||||
|
HUNDREDS = ["", "сто", "двести", "триста", "четыреста", "пятьсот", "шестьсот", "семьсот", "восемьсот", "девятьсот"]
|
||||||
|
SCALES = [
|
||||||
|
None,
|
||||||
|
("тысяча", "тысячи", "тысяч"),
|
||||||
|
("миллион", "миллиона", "миллионов"),
|
||||||
|
("миллиард", "миллиарда", "миллиардов"),
|
||||||
|
("триллион", "триллиона", "триллионов"),
|
||||||
|
("квадриллион", "квадриллиона", "квадриллионов"),
|
||||||
|
]
|
||||||
|
NUMBER_PATTERN = re.compile(r"(?<![^\W_])([+-]?)(\d{1,18})(?:[,.](\d+))?(?![^\W_])", re.UNICODE)
|
||||||
|
YEAR_IN_PREPOSITIONAL_PATTERN = re.compile(r"(?<![^\W_])(\d{4})\s+(году|г\.)", re.IGNORECASE | re.UNICODE)
|
||||||
|
|
||||||
|
SOFT_LETTERS = set("яёюиье")
|
||||||
|
SYLLABLE_STARTS = set("#ъьаёуюэеиы-")
|
||||||
|
SOFT_HARD_CONSONANTS = {
|
||||||
|
"б": "b", "в": "v", "г": "g", "д": "d", "з": "z", "к": "k", "л": "l",
|
||||||
|
"м": "m", "н": "n", "п": "p", "р": "r", "с": "s", "т": "t", "ф": "f", "х": "h",
|
||||||
|
}
|
||||||
|
OTHER_CONSONANTS = {"ж": "zh", "ц": "c", "ч": "ch", "ш": "sh", "щ": "sch", "й": "j"}
|
||||||
|
VOWELS = {"а": "a", "я": "a", "у": "u", "ю": "u", "о": "o", "ё": "o", "э": "e", "е": "e", "и": "i", "ы": "y"}
|
||||||
|
|
||||||
|
|
||||||
|
def scale_form(value: int, forms: tuple[str, str, str]) -> str:
|
||||||
|
if 11 <= value % 100 <= 14:
|
||||||
|
return forms[2]
|
||||||
|
return forms[0] if value % 10 == 1 else forms[1] if value % 10 in (2, 3, 4) else forms[2]
|
||||||
|
|
||||||
|
|
||||||
|
def under_thousand_to_words(value: int, feminine: bool) -> str:
|
||||||
|
words: list[str] = []
|
||||||
|
if value // 100:
|
||||||
|
words.append(HUNDREDS[value // 100])
|
||||||
|
tail = value % 100
|
||||||
|
if 10 <= tail <= 19:
|
||||||
|
words.append(TEENS[tail - 10])
|
||||||
|
else:
|
||||||
|
if tail // 10:
|
||||||
|
words.append(TENS[tail // 10])
|
||||||
|
units = tail % 10
|
||||||
|
if units:
|
||||||
|
words.append("одна" if feminine and units == 1 else "две" if feminine and units == 2 else DIGITS[units])
|
||||||
|
return " ".join(words)
|
||||||
|
|
||||||
|
|
||||||
|
def integer_to_words(value: int) -> str:
|
||||||
|
if value == 0:
|
||||||
|
return DIGITS[0]
|
||||||
|
groups: list[int] = []
|
||||||
|
while value:
|
||||||
|
groups.append(value % 1_000)
|
||||||
|
value //= 1_000
|
||||||
|
result: list[str] = []
|
||||||
|
for group_index in range(len(groups) - 1, -1, -1):
|
||||||
|
group = groups[group_index]
|
||||||
|
if not group:
|
||||||
|
continue
|
||||||
|
words = under_thousand_to_words(group, group_index == 1)
|
||||||
|
forms = SCALES[group_index] if group_index < len(SCALES) else None
|
||||||
|
result.append(words if forms is None else f"{words} {scale_form(group, forms)}")
|
||||||
|
return " ".join(result)
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_numbers(text: str) -> str:
|
||||||
|
prepared = re.sub(r"(?<=\d)[\u00a0\u202f](?=\d{3}(?:\D|$))", "", text)
|
||||||
|
prepared = re.sub(r"№\s*(?=\d)", "номер ", prepared)
|
||||||
|
|
||||||
|
def replace(match: re.Match[str]) -> str:
|
||||||
|
sign = "минус " if match.group(1) == "-" else "плюс " if match.group(1) == "+" else ""
|
||||||
|
integer_digits = match.group(2)
|
||||||
|
if len(integer_digits) > 1 and integer_digits.startswith("0"):
|
||||||
|
integer = " ".join(DIGITS[int(char)] for char in integer_digits)
|
||||||
|
else:
|
||||||
|
integer = integer_to_words(int(integer_digits))
|
||||||
|
fraction = match.group(3)
|
||||||
|
return f"{sign}{integer}" if not fraction else f"{sign}{integer} запятая {' '.join(DIGITS[int(char)] for char in fraction)}"
|
||||||
|
|
||||||
|
prepared = YEAR_IN_PREPOSITIONAL_PATTERN.sub(
|
||||||
|
lambda match: f"{year_in_prepositional_words(int(match.group(1)))} {match.group(2)}",
|
||||||
|
prepared,
|
||||||
|
)
|
||||||
|
return NUMBER_PATTERN.sub(replace, prepared)
|
||||||
|
|
||||||
|
|
||||||
|
ORDINAL_UNITS_PREPOSITIONAL = ("", "первом", "втором", "третьем", "четвёртом", "пятом", "шестом", "седьмом", "восьмом", "девятом")
|
||||||
|
ORDINAL_TEENS_PREPOSITIONAL = ("десятом", "одиннадцатом", "двенадцатом", "тринадцатом", "четырнадцатом", "пятнадцатом", "шестнадцатом", "семнадцатом", "восемнадцатом", "девятнадцатом")
|
||||||
|
ORDINAL_TENS_PREPOSITIONAL = ("", "десятом", "двадцатом", "тридцатом", "сороковом", "пятидесятом", "шестидесятом", "семидесятом", "восьмидесятом", "девяностом")
|
||||||
|
ORDINAL_HUNDREDS_PREPOSITIONAL = ("", "сотом", "двухсотом", "трёхсотом", "четырёхсотом", "пятисотом", "шестисотом", "семисотом", "восьмисотом", "девятисотом")
|
||||||
|
EXACT_THOUSANDTH_PREPOSITIONAL = ("", "тысячном", "двухтысячном", "трёхтысячном", "четырёхтысячном", "пятитысячном", "шеститысячном", "семитысячном", "восьмитысячном", "девятитысячном")
|
||||||
|
|
||||||
|
|
||||||
|
def ordinal_under_thousand_prepositional(value: int) -> str:
|
||||||
|
last_two = value % 100
|
||||||
|
units = value % 10
|
||||||
|
if units and not 11 <= last_two <= 19:
|
||||||
|
return f"{integer_to_words(value - units) if value > units else ''} {ORDINAL_UNITS_PREPOSITIONAL[units]}".strip()
|
||||||
|
if 11 <= last_two <= 19:
|
||||||
|
return f"{integer_to_words(value - last_two) if value > last_two else ''} {ORDINAL_TEENS_PREPOSITIONAL[last_two - 10]}".strip()
|
||||||
|
if last_two >= 20:
|
||||||
|
return f"{integer_to_words(value - last_two) if value > last_two else ''} {ORDINAL_TENS_PREPOSITIONAL[last_two // 10]}".strip()
|
||||||
|
return ORDINAL_HUNDREDS_PREPOSITIONAL[value // 100]
|
||||||
|
|
||||||
|
|
||||||
|
def year_in_prepositional_words(value: int) -> str:
|
||||||
|
if not 1000 <= value <= 9999:
|
||||||
|
raise ValueError(value)
|
||||||
|
if value % 1000 == 0:
|
||||||
|
return EXACT_THOUSANDTH_PREPOSITIONAL[value // 1000]
|
||||||
|
thousands = value // 1000
|
||||||
|
thousands_words = under_thousand_to_words(thousands, True)
|
||||||
|
thousands_prefix = "" if thousands_words == "одна" else thousands_words + " "
|
||||||
|
return f"{thousands_prefix}{scale_form(thousands, SCALES[1])} {ordinal_under_thousand_prepositional(value % 1000)}"
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_text(text: str) -> str:
|
||||||
|
stressed = re.sub(r"([А-Яа-яЁё])\u0301", lambda match: "+" + match.group(1), text)
|
||||||
|
translated = normalize_numbers(stressed).lower().replace("—", "-").replace("–", "-").replace("…", "...")
|
||||||
|
translated = translated.translate(str.maketrans({"«": '"', "»": '"', "“": '"', "”": '"'}))
|
||||||
|
return re.sub(r"\s+", " ", translated).strip()
|
||||||
|
|
||||||
|
|
||||||
|
class WordPieceTokenizer:
|
||||||
|
def __init__(self, path: Path):
|
||||||
|
self.tokens = path.read_text(encoding="utf-8").splitlines()
|
||||||
|
self.vocabulary = {token: index for index, token in enumerate(self.tokens)}
|
||||||
|
for special in ("[CLS]", "[SEP]", "[UNK]"):
|
||||||
|
if special not in self.vocabulary:
|
||||||
|
raise ValueError(f"Missing {special} in {path}")
|
||||||
|
|
||||||
|
def encode(self, text: str) -> tuple[np.ndarray, list[int]]:
|
||||||
|
ids = [self.vocabulary["[CLS]"]]
|
||||||
|
positions = [0]
|
||||||
|
for token in self._basic_tokens(text):
|
||||||
|
punctuation = token in BERT_EXCLUDED_PUNCTUATION
|
||||||
|
pieces = [self.vocabulary.get(token, self.vocabulary["[UNK]"])] if punctuation else self._word_pieces(token)
|
||||||
|
if not punctuation:
|
||||||
|
positions.append(len(ids))
|
||||||
|
ids.extend(pieces)
|
||||||
|
ids.append(self.vocabulary["[SEP]"])
|
||||||
|
positions.append(len(ids) - 1)
|
||||||
|
return np.asarray(ids, dtype=np.int64), positions
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _basic_tokens(text: str) -> list[str]:
|
||||||
|
result: list[str] = []
|
||||||
|
word: list[str] = []
|
||||||
|
for char in text:
|
||||||
|
if char.isspace():
|
||||||
|
if word:
|
||||||
|
result.append("".join(word)); word.clear()
|
||||||
|
elif char.isalnum():
|
||||||
|
word.append(char.lower())
|
||||||
|
elif char == "+":
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
if word:
|
||||||
|
result.append("".join(word)); word.clear()
|
||||||
|
result.append(char)
|
||||||
|
if word:
|
||||||
|
result.append("".join(word))
|
||||||
|
return result
|
||||||
|
|
||||||
|
def _word_pieces(self, token: str) -> list[int]:
|
||||||
|
if len(token) > MAX_WORD_PIECE_CHARS:
|
||||||
|
return [self.vocabulary["[UNK]"]]
|
||||||
|
result: list[int] = []
|
||||||
|
start = 0
|
||||||
|
while start < len(token):
|
||||||
|
found: int | None = None
|
||||||
|
found_end = start
|
||||||
|
for end in range(len(token), start, -1):
|
||||||
|
piece = token[start:end] if start == 0 else f"##{token[start:end]}"
|
||||||
|
if piece in self.vocabulary:
|
||||||
|
found = self.vocabulary[piece]; found_end = end; break
|
||||||
|
if found is None:
|
||||||
|
return [self.vocabulary["[UNK]"]]
|
||||||
|
result.append(found); start = found_end
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
class PronunciationDictionary:
|
||||||
|
def __init__(self, dictionary_path: Path, index_path: Path):
|
||||||
|
self.source = dictionary_path.open("rb")
|
||||||
|
entries = []
|
||||||
|
for line in index_path.read_text(encoding="utf-8").splitlines():
|
||||||
|
word, offset = line.rsplit("\t", 1)
|
||||||
|
entries.append((word, int(offset)))
|
||||||
|
self.words = [entry[0] for entry in entries]
|
||||||
|
self.offsets = [entry[1] for entry in entries]
|
||||||
|
self.length = dictionary_path.stat().st_size
|
||||||
|
self.cache: OrderedDict[str, str | None] = OrderedDict()
|
||||||
|
|
||||||
|
def close(self) -> None:
|
||||||
|
self.source.close()
|
||||||
|
|
||||||
|
def find(self, word: str) -> str | None:
|
||||||
|
if word in self.cache:
|
||||||
|
value = self.cache.pop(word); self.cache[word] = value; return value
|
||||||
|
index = max(0, bisect.bisect_right(self.words, word) - 1)
|
||||||
|
end = self.offsets[index + 1] if index + 1 < len(self.offsets) else self.length
|
||||||
|
self.source.seek(self.offsets[index])
|
||||||
|
value: str | None = None
|
||||||
|
while self.source.tell() < end:
|
||||||
|
raw = self.source.readline()
|
||||||
|
if not raw:
|
||||||
|
break
|
||||||
|
entry_word, _, pronunciation = raw.decode("utf-8").rstrip("\r\n").partition("\t")
|
||||||
|
if entry_word < word:
|
||||||
|
continue
|
||||||
|
if entry_word == word:
|
||||||
|
value = pronunciation.strip() or None
|
||||||
|
break
|
||||||
|
self.cache[word] = value
|
||||||
|
if len(self.cache) > CACHE_SIZE:
|
||||||
|
self.cache.popitem(last=False)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def russian_g2p(word: str) -> str:
|
||||||
|
marked: list[tuple[str, int]] = []
|
||||||
|
stress = 0
|
||||||
|
for char in f"#{word}#":
|
||||||
|
if char == "+":
|
||||||
|
stress = 1
|
||||||
|
else:
|
||||||
|
marked.append((char, stress)); stress = 0
|
||||||
|
phones: list[str] = []
|
||||||
|
for index, (char, accent) in enumerate(marked):
|
||||||
|
previous = marked[index - 1][0] if index else "#"
|
||||||
|
following = marked[index + 1][0] if index + 1 < len(marked) else None
|
||||||
|
if char in SOFT_HARD_CONSONANTS:
|
||||||
|
phones.append(SOFT_HARD_CONSONANTS[char] + ("j" if following in SOFT_LETTERS else ""))
|
||||||
|
elif char in OTHER_CONSONANTS:
|
||||||
|
phones.append(OTHER_CONSONANTS[char])
|
||||||
|
elif char in VOWELS:
|
||||||
|
if previous in SYLLABLE_STARTS and char in set("яюеё"):
|
||||||
|
phones.append("j")
|
||||||
|
phones.append(f"{VOWELS[char]}{accent}")
|
||||||
|
return " ".join(phones)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class RawPhone:
|
||||||
|
phone: str
|
||||||
|
punctuation: list[str]
|
||||||
|
in_quote: int
|
||||||
|
bert_word_index: int
|
||||||
|
|
||||||
|
|
||||||
|
def build_multistream_phones(text: str, dictionary: PronunciationDictionary) -> tuple[np.ndarray, list[int]]:
|
||||||
|
raw = [RawPhone("^", [], 0, 0)]
|
||||||
|
word: list[str] = []
|
||||||
|
pending: list[str] = []
|
||||||
|
in_quote = 0
|
||||||
|
bert_word_index = 1
|
||||||
|
|
||||||
|
def flush_word() -> None:
|
||||||
|
nonlocal bert_word_index
|
||||||
|
if not word:
|
||||||
|
return
|
||||||
|
value = "".join(word)
|
||||||
|
pronunciation = dictionary.find(value) or russian_g2p(value)
|
||||||
|
raw.extend(RawPhone(phone, [], in_quote, bert_word_index) for phone in pronunciation.split() if phone)
|
||||||
|
word.clear(); bert_word_index += 1
|
||||||
|
|
||||||
|
def append_space() -> None:
|
||||||
|
raw.append(RawPhone(" ", list(pending), in_quote, bert_word_index)); pending.clear()
|
||||||
|
|
||||||
|
index = 0
|
||||||
|
while index < len(text):
|
||||||
|
if text.startswith("...", index):
|
||||||
|
flush_word(); pending.append("..."); index += 3; continue
|
||||||
|
char = text[index]
|
||||||
|
if char in {'"', "«", "»", "“", "”"}:
|
||||||
|
flush_word(); in_quote = 1 - in_quote
|
||||||
|
elif char.isspace():
|
||||||
|
flush_word(); append_space()
|
||||||
|
elif char == "-" and index > 0 and index + 1 < len(text) and text[index - 1].isalnum() and text[index + 1].isalnum():
|
||||||
|
flush_word()
|
||||||
|
elif char in MULTISTREAM_PUNCTUATION:
|
||||||
|
flush_word(); pending.append(char)
|
||||||
|
elif char.isalnum() or char == "+":
|
||||||
|
word.append(char)
|
||||||
|
index += 1
|
||||||
|
flush_word(); append_space(); raw.append(RawPhone("$", [], 0, bert_word_index))
|
||||||
|
|
||||||
|
last_punctuation = " "
|
||||||
|
last_sentence_punctuation = " "
|
||||||
|
reversed_features: list[list[int]] = []
|
||||||
|
reversed_positions: list[int] = []
|
||||||
|
for phone in reversed(raw):
|
||||||
|
for candidate in ("...", ".", "!", "?", "-"):
|
||||||
|
if candidate in phone.punctuation:
|
||||||
|
last_sentence_punctuation = candidate; break
|
||||||
|
if phone.punctuation:
|
||||||
|
last_punctuation = phone.punctuation[0]
|
||||||
|
current = phone.punctuation[0] if phone.punctuation else "_"
|
||||||
|
values = [PHONEME_IDS[phone.phone], PHONEME_IDS[current], phone.in_quote, PHONEME_IDS[last_punctuation], PHONEME_IDS[last_sentence_punctuation]]
|
||||||
|
reversed_features.append(values); reversed_positions.append(phone.bert_word_index)
|
||||||
|
features = np.asarray(list(reversed(reversed_features)), dtype=np.int64).T[np.newaxis, :, :]
|
||||||
|
return features, list(reversed(reversed_positions))
|
||||||
|
|
||||||
|
|
||||||
|
class Frontend:
|
||||||
|
def __init__(self, assets: Path, bert_model: Path | None = None):
|
||||||
|
import onnxruntime as ort
|
||||||
|
self.tokenizer = WordPieceTokenizer(assets / "vocab.txt")
|
||||||
|
self.dictionary = PronunciationDictionary(assets / "dict.tsv", assets / "dictionary.index")
|
||||||
|
self.bert = ort.InferenceSession(str(bert_model or assets / "bert.int8.onnx"), providers=["CPUExecutionProvider"])
|
||||||
|
|
||||||
|
def close(self) -> None:
|
||||||
|
self.dictionary.close()
|
||||||
|
|
||||||
|
def prepare(self, text: str, articulation: float = 1.0) -> dict[str, np.ndarray]:
|
||||||
|
normalized = normalize_text(text)
|
||||||
|
ids, embedding_positions = self.tokenizer.encode(normalized)
|
||||||
|
shape = (1, ids.size)
|
||||||
|
token_embeddings = self.bert.run(None, {
|
||||||
|
"input_ids": ids.reshape(shape),
|
||||||
|
"attention_mask": np.ones(shape, dtype=np.int64),
|
||||||
|
"token_type_ids": np.zeros(shape, dtype=np.int64),
|
||||||
|
})[0]
|
||||||
|
selected = token_embeddings[np.asarray(embedding_positions, dtype=np.int64)]
|
||||||
|
features, phone_positions = build_multistream_phones(normalized, self.dictionary)
|
||||||
|
phone_embeddings = selected[np.clip(np.asarray(phone_positions), 0, len(selected) - 1)]
|
||||||
|
bert = phone_embeddings.T[np.newaxis, :, :].astype(np.float32, copy=False)
|
||||||
|
time_steps = features.shape[2]
|
||||||
|
if bert.shape != (1, BERT_DIMENSIONS, time_steps):
|
||||||
|
raise ValueError(f"Unexpected BERT shape {bert.shape} for {text!r}")
|
||||||
|
return {
|
||||||
|
"input": features,
|
||||||
|
"input_lengths": np.asarray([time_steps], dtype=np.int64),
|
||||||
|
"scales": np.asarray([0.8, articulation, 0.8], dtype=np.float32),
|
||||||
|
"sid": np.asarray([SPEAKER_ID], dtype=np.int64),
|
||||||
|
"bert": bert,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def load_corpus(path: Path) -> list[str]:
|
||||||
|
lines = [line.strip() for line in path.read_text(encoding="utf-8-sig").splitlines()]
|
||||||
|
result = [line for line in lines if line and not line.startswith("#")]
|
||||||
|
if not result:
|
||||||
|
raise ValueError(f"No phrases in {path}")
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def command_prepare(args: argparse.Namespace) -> None:
|
||||||
|
args.output.mkdir(parents=True, exist_ok=False)
|
||||||
|
frontend = Frontend(args.assets, args.bert_model)
|
||||||
|
manifest = []
|
||||||
|
try:
|
||||||
|
for index, text in enumerate(load_corpus(args.corpus)):
|
||||||
|
feeds = frontend.prepare(text, args.articulation)
|
||||||
|
name = f"sample-{index:04d}.npz"
|
||||||
|
np.savez_compressed(args.output / name, **feeds)
|
||||||
|
manifest.append({"file": name, "text": text, "time_steps": int(feeds["input_lengths"][0])})
|
||||||
|
print(f"prepared {index + 1}: phones={manifest[-1]['time_steps']} text={text}")
|
||||||
|
finally:
|
||||||
|
frontend.close()
|
||||||
|
(args.output / "manifest.json").write_text(json.dumps(manifest, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||||
|
print(f"samples={len(manifest)} output={args.output}")
|
||||||
|
|
||||||
|
|
||||||
|
class NpzCalibrationReader:
|
||||||
|
def __init__(self, directory: Path):
|
||||||
|
self.files = sorted(directory.glob("sample-*.npz"))
|
||||||
|
self.iterator: Iterator[Path] | None = None
|
||||||
|
if not self.files:
|
||||||
|
raise ValueError(f"No sample-*.npz in {directory}")
|
||||||
|
|
||||||
|
def get_next(self) -> dict[str, np.ndarray] | None:
|
||||||
|
if self.iterator is None:
|
||||||
|
self.iterator = iter(self.files)
|
||||||
|
try:
|
||||||
|
path = next(self.iterator)
|
||||||
|
except StopIteration:
|
||||||
|
return None
|
||||||
|
with np.load(path) as data:
|
||||||
|
return {name: data[name] for name in data.files}
|
||||||
|
|
||||||
|
def rewind(self) -> None:
|
||||||
|
self.iterator = iter(self.files)
|
||||||
|
|
||||||
|
|
||||||
|
def command_quantize(args: argparse.Namespace) -> None:
|
||||||
|
from onnxruntime.quantization import QuantFormat, QuantType, quantize_static
|
||||||
|
if args.output.exists():
|
||||||
|
raise FileExistsError(f"Refusing to overwrite {args.output}")
|
||||||
|
started = time.time()
|
||||||
|
quantize_static(
|
||||||
|
str(args.model), str(args.output), NpzCalibrationReader(args.calibration),
|
||||||
|
quant_format=QuantFormat.QDQ, op_types_to_quantize=["Conv"], per_channel=True,
|
||||||
|
activation_type=QuantType.QInt8, weight_type=QuantType.QInt8,
|
||||||
|
)
|
||||||
|
source_size = args.model.stat().st_size
|
||||||
|
target_size = args.output.stat().st_size
|
||||||
|
print(json.dumps({
|
||||||
|
"source_bytes": source_size, "target_bytes": target_size,
|
||||||
|
"saving_bytes": source_size - target_size,
|
||||||
|
"saving_percent": round((source_size - target_size) * 100 / source_size, 4),
|
||||||
|
"seconds": round(time.time() - started, 3),
|
||||||
|
}, indent=2))
|
||||||
|
|
||||||
|
|
||||||
|
def write_wav(path: Path, samples: np.ndarray) -> None:
|
||||||
|
pcm = (np.clip(samples.reshape(-1), -1.0, 1.0) * 32767.0).astype("<i2").tobytes()
|
||||||
|
with wave.open(str(path), "wb") as output:
|
||||||
|
output.setnchannels(1); output.setsampwidth(2); output.setframerate(SAMPLE_RATE); output.writeframes(pcm)
|
||||||
|
|
||||||
|
|
||||||
|
def command_compare(args: argparse.Namespace) -> None:
|
||||||
|
import onnxruntime as ort
|
||||||
|
args.output.mkdir(parents=True, exist_ok=False)
|
||||||
|
options = ort.SessionOptions(); options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_DISABLE_ALL
|
||||||
|
baseline = ort.InferenceSession(str(args.baseline), sess_options=options, providers=["CPUExecutionProvider"])
|
||||||
|
candidate = ort.InferenceSession(str(args.candidate), sess_options=options, providers=["CPUExecutionProvider"])
|
||||||
|
manifest = json.loads((args.calibration / "manifest.json").read_text(encoding="utf-8"))
|
||||||
|
report = []
|
||||||
|
for item in manifest[: args.limit]:
|
||||||
|
with np.load(args.calibration / item["file"]) as data:
|
||||||
|
feeds = {name: data[name] for name in data.files}
|
||||||
|
feeds["scales"] = np.asarray([0.0, float(feeds["scales"][1]), 0.0], dtype=np.float32)
|
||||||
|
base_wav, base_length = baseline.run(None, feeds)
|
||||||
|
test_wav, test_length = candidate.run(None, feeds)
|
||||||
|
common = min(base_wav.size, test_wav.size)
|
||||||
|
difference = base_wav.reshape(-1)[:common] - test_wav.reshape(-1)[:common]
|
||||||
|
entry = {
|
||||||
|
"text": item["text"], "baseline_samples": int(base_wav.size), "candidate_samples": int(test_wav.size),
|
||||||
|
"baseline_length": np.asarray(base_length).tolist(), "candidate_length": np.asarray(test_length).tolist(),
|
||||||
|
"mae": float(np.mean(np.abs(difference))), "rmse": float(math.sqrt(np.mean(difference * difference))),
|
||||||
|
}
|
||||||
|
report.append(entry)
|
||||||
|
number = len(report)
|
||||||
|
write_wav(args.output / f"{number:02d}-baseline.wav", base_wav)
|
||||||
|
write_wav(args.output / f"{number:02d}-candidate.wav", test_wav)
|
||||||
|
print(f"compared {number}: mae={entry['mae']:.6f} rmse={entry['rmse']:.6f}")
|
||||||
|
(args.output / "report.json").write_text(json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||||
|
|
||||||
|
|
||||||
|
def command_compare_feeds(args: argparse.Namespace) -> None:
|
||||||
|
import onnxruntime as ort
|
||||||
|
args.output.mkdir(parents=True, exist_ok=args.resume)
|
||||||
|
options = ort.SessionOptions(); options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_DISABLE_ALL
|
||||||
|
acoustic = ort.InferenceSession(str(args.model), sess_options=options, providers=["CPUExecutionProvider"])
|
||||||
|
baseline_manifest = json.loads((args.baseline_calibration / "manifest.json").read_text(encoding="utf-8"))
|
||||||
|
candidate_manifest = json.loads((args.candidate_calibration / "manifest.json").read_text(encoding="utf-8"))
|
||||||
|
if [item["text"] for item in baseline_manifest] != [item["text"] for item in candidate_manifest]:
|
||||||
|
raise ValueError("Baseline and candidate manifests contain different texts")
|
||||||
|
report = []
|
||||||
|
for baseline_item, candidate_item in zip(baseline_manifest[: args.limit or None], candidate_manifest[: args.limit or None]):
|
||||||
|
with np.load(args.baseline_calibration / baseline_item["file"]) as data:
|
||||||
|
baseline_feeds = {name: data[name] for name in data.files}
|
||||||
|
with np.load(args.candidate_calibration / candidate_item["file"]) as data:
|
||||||
|
candidate_feeds = {name: data[name] for name in data.files}
|
||||||
|
baseline_feeds["scales"] = np.asarray([0.0, float(baseline_feeds["scales"][1]), 0.0], dtype=np.float32)
|
||||||
|
candidate_feeds["scales"] = np.asarray([0.0, float(candidate_feeds["scales"][1]), 0.0], dtype=np.float32)
|
||||||
|
baseline_wav, baseline_length = acoustic.run(None, baseline_feeds)
|
||||||
|
candidate_wav, candidate_length = acoustic.run(None, candidate_feeds)
|
||||||
|
common = min(baseline_wav.size, candidate_wav.size)
|
||||||
|
difference = baseline_wav.reshape(-1)[:common] - candidate_wav.reshape(-1)[:common]
|
||||||
|
entry = {
|
||||||
|
"text": baseline_item["text"],
|
||||||
|
"baseline_samples": int(baseline_wav.size), "candidate_samples": int(candidate_wav.size),
|
||||||
|
"baseline_length": np.asarray(baseline_length).tolist(),
|
||||||
|
"candidate_length": np.asarray(candidate_length).tolist(),
|
||||||
|
"mae": float(np.mean(np.abs(difference))),
|
||||||
|
"rmse": float(math.sqrt(np.mean(difference * difference))),
|
||||||
|
}
|
||||||
|
report.append(entry)
|
||||||
|
number = len(report)
|
||||||
|
baseline_path = args.output / f"{number:03d}-baseline.wav"
|
||||||
|
candidate_path = args.output / f"{number:03d}-student-int8.wav"
|
||||||
|
if not baseline_path.exists():
|
||||||
|
write_wav(baseline_path, baseline_wav)
|
||||||
|
if not candidate_path.exists():
|
||||||
|
write_wav(candidate_path, candidate_wav)
|
||||||
|
temporary_report = args.output / "report.json.tmp"
|
||||||
|
temporary_report.write_text(json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||||
|
temporary_report.replace(args.output / "report.json")
|
||||||
|
print(f"rendered {number}: mae={entry['mae']:.6f} rmse={entry['rmse']:.6f}", flush=True)
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
commands = parser.add_subparsers(dest="command", required=True)
|
||||||
|
prepare = commands.add_parser("prepare")
|
||||||
|
prepare.add_argument("--assets", type=Path, required=True)
|
||||||
|
prepare.add_argument("--corpus", type=Path, required=True)
|
||||||
|
prepare.add_argument("--output", type=Path, required=True)
|
||||||
|
prepare.add_argument("--articulation", type=float, default=1.0)
|
||||||
|
prepare.add_argument("--bert-model", type=Path)
|
||||||
|
prepare.set_defaults(handler=command_prepare)
|
||||||
|
quantize = commands.add_parser("quantize")
|
||||||
|
quantize.add_argument("--model", type=Path, required=True)
|
||||||
|
quantize.add_argument("--calibration", type=Path, required=True)
|
||||||
|
quantize.add_argument("--output", type=Path, required=True)
|
||||||
|
quantize.set_defaults(handler=command_quantize)
|
||||||
|
compare = commands.add_parser("compare")
|
||||||
|
compare.add_argument("--baseline", type=Path, required=True)
|
||||||
|
compare.add_argument("--candidate", type=Path, required=True)
|
||||||
|
compare.add_argument("--calibration", type=Path, required=True)
|
||||||
|
compare.add_argument("--output", type=Path, required=True)
|
||||||
|
compare.add_argument("--limit", type=int, default=3)
|
||||||
|
compare.set_defaults(handler=command_compare)
|
||||||
|
compare_feeds = commands.add_parser("compare-feeds")
|
||||||
|
compare_feeds.add_argument("--model", type=Path, required=True)
|
||||||
|
compare_feeds.add_argument("--baseline-calibration", type=Path, required=True)
|
||||||
|
compare_feeds.add_argument("--candidate-calibration", type=Path, required=True)
|
||||||
|
compare_feeds.add_argument("--output", type=Path, required=True)
|
||||||
|
compare_feeds.add_argument("--limit", type=int, default=0)
|
||||||
|
compare_feeds.add_argument("--resume", action="store_true")
|
||||||
|
compare_feeds.set_defaults(handler=command_compare_feeds)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> int:
|
||||||
|
args = parse_args()
|
||||||
|
args.handler(args)
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(main())
|
||||||
Reference in New Issue
Block a user