mirror of
https://github.com/cmusphinx/sphinxtrain.git
synced 2026-06-16 13:14:30 +00:00
vastly elaborate on the s3dict module
git-svn-id: svn+ssh://svn.code.sf.net/p/cmusphinx/code/trunk/SphinxTrain@9798 94700074-3cef-4d97-a70e-9c8c206c02f5
This commit is contained in:
+138
-30
@@ -12,6 +12,7 @@ SphinxTrain, Sphinx-III, and PocketSphinx.
|
||||
__author__ = "David Huggins-Daines <dhuggins@cs.cmu.edu>"
|
||||
__version__ = "$Revision $"
|
||||
|
||||
from collections import defaultdict
|
||||
import re
|
||||
|
||||
def open(file):
|
||||
@@ -21,11 +22,36 @@ class S3Dict(dict):
|
||||
"""
|
||||
Class for reading / processing Sphinx format dictionary files.
|
||||
"""
|
||||
def __init__(self, infile=None):
|
||||
self.phoneset = set()
|
||||
def __init__(self, infile=None, preserve_alts=False):
|
||||
self.preserve_alts = preserve_alts
|
||||
self.phoneset = defaultdict(int)
|
||||
self.maxalt = defaultdict(int)
|
||||
if infile != None:
|
||||
self.read(infile)
|
||||
|
||||
def __getitem__(self, key):
|
||||
if not isinstance(key, tuple):
|
||||
m = self.altre.match(key)
|
||||
if m:
|
||||
word, alt = m.groups()
|
||||
return self.get_alt_phones(word, int(alt))
|
||||
else:
|
||||
return self.get_phones(key)
|
||||
else:
|
||||
return self.get_alt_phones(*key)
|
||||
|
||||
def __putitem__(self, key, val):
|
||||
if not isinstance(key, tuple):
|
||||
m = self.altre.match(key)
|
||||
if m:
|
||||
word, alt = m.groups()
|
||||
return self.set_alt_phones(word, int(alt), val)
|
||||
else:
|
||||
self.set_phones(key, val)
|
||||
else:
|
||||
w, p = key
|
||||
self.set_alt_phones(w, p, val)
|
||||
|
||||
altre = re.compile(r'(.*)\(([^\)]+)\)')
|
||||
def read(self, infile):
|
||||
"""
|
||||
@@ -45,14 +71,8 @@ class S3Dict(dict):
|
||||
spam = line.split()
|
||||
word = unicode(spam[0], 'utf8')
|
||||
phones = spam[1:]
|
||||
m = self.altre.match(word)
|
||||
if m:
|
||||
word, alt = m.groups()
|
||||
self[word,int(alt)] = phones
|
||||
else:
|
||||
self[word,1] = phones
|
||||
for ph in phones:
|
||||
self.phoneset.add(ph)
|
||||
# Why can't we just say self[word] = phones?
|
||||
self.__putitem__(word, phones)
|
||||
|
||||
def write(self, outfile):
|
||||
"""
|
||||
@@ -70,24 +90,112 @@ class S3Dict(dict):
|
||||
word = "%s(%d)" % (word, alt)
|
||||
fh.write("%-30s %s\n" % (word, " ".join(self[k])))
|
||||
|
||||
def __getitem__(self, key):
|
||||
if not isinstance(key, tuple):
|
||||
m = self.altre.match(key)
|
||||
if m:
|
||||
word, alt = m.groups()
|
||||
return self.get((word,int(alt)))
|
||||
else:
|
||||
return self.get((key,1))
|
||||
else:
|
||||
return self.get(key)
|
||||
def get_phones(self, word):
|
||||
"""
|
||||
Get default pronunciation for word.
|
||||
|
||||
def getalts(self, key):
|
||||
i = 1
|
||||
alts = []
|
||||
while True:
|
||||
if (key,i) in self:
|
||||
alts.append(self.get((key,i)))
|
||||
else:
|
||||
break
|
||||
i += 1
|
||||
return alts
|
||||
If word is not present, KeyError will be raised.
|
||||
"""
|
||||
return dict.__getitem__(self, (word, 1))
|
||||
|
||||
def get_alt_phones(self, word, alt):
|
||||
"""
|
||||
Get alternate pronunciaition #alt for word.
|
||||
|
||||
Alternate pronunciations are numbered from 1, where 1 is the
|
||||
default pronunciation. If word is not present, KeyError will
|
||||
be raised. If no alternate pronunciation alt exists,
|
||||
IndexError will be raised.
|
||||
"""
|
||||
|
||||
if (word, 1) not in self:
|
||||
raise KeyError
|
||||
elif (word,alt) not in self:
|
||||
raise IndexError, "Alternate pronunciation index %d does not exist" % alt
|
||||
return dict.__getitem__(self, (word, alt))
|
||||
|
||||
def set_phones(self, word, phones):
|
||||
"""
|
||||
Set default pronunciation for word.
|
||||
"""
|
||||
dict.__setitem__(self, (word, 1), phones)
|
||||
self.maxalt[word] = 1
|
||||
for ph in phones:
|
||||
self.phoneset[ph] += 1 # FIXME: should make a class for this
|
||||
|
||||
def set_alt_phones(self, word, alt, phones):
|
||||
"""
|
||||
Set alternate pronunciaition #alt for word.
|
||||
|
||||
If alt is greater than the maximum alternate pronunciation
|
||||
index plus one for this dictionary, IndexError will be raised.
|
||||
|
||||
"""
|
||||
if alt > self.maxalt[word] + 1:
|
||||
raise IndexError, "Alternate pronunciation index %d too high" % alt
|
||||
dict.__setitem__(self, (word, alt), phones)
|
||||
self.maxalt[word] = max(alt, self.maxalt[word])
|
||||
for ph in phones:
|
||||
self.phoneset[ph] += 1
|
||||
|
||||
def add_alt_phones(self, word, phones):
|
||||
"""
|
||||
Add a new alternate pronunciation for word.
|
||||
"""
|
||||
dict.__setitem__(self, (word, self.maxalt[word] + 1), phones)
|
||||
self.maxalt[word] += 1
|
||||
for ph in phones:
|
||||
self.phoneset[ph] += 1
|
||||
|
||||
def del_alt_phones(self, word, alt):
|
||||
"""
|
||||
Delete alternate pronunciation alt for word.
|
||||
|
||||
If no such alternate pronunciation exists, IndexError will be
|
||||
raised. If this S3Dict was created with preserve_alts=True,
|
||||
the indices of remaining alternate pronunciations will be
|
||||
preserved (you can still this alternative index with
|
||||
set_alt_phones()). Otherwise, the remaining alternate
|
||||
pronunciations will be renumbered accordingly.
|
||||
|
||||
"""
|
||||
if (word, alt) not in self:
|
||||
raise IndexError, "Alternate pronunciation index %d does not exist" % alt
|
||||
for ph in self[word, alt]:
|
||||
self.phoneset[ph] -= 1
|
||||
if self.phoneset[ph] == 0: # FIXME: make a class
|
||||
del self.phoneset[ph]
|
||||
del self[word, alt]
|
||||
if alt == self.maxalt[word]:
|
||||
self.maxalt[word] -= 1
|
||||
if not self.preserve_alts:
|
||||
alts = list(self.alts(word))
|
||||
self.del_phones(word)
|
||||
for i, phones in enumerate(alts):
|
||||
dict.__setitem__(self, (word, i + 1), phones)
|
||||
self.maxalt[word] = len(alts)
|
||||
|
||||
def del_phones(self, word):
|
||||
"""
|
||||
Delete all pronunciations for word.
|
||||
|
||||
If you only wish to delete the default pronunciation (it is
|
||||
strongly suggested that you don't do this), use
|
||||
del_alt_phones(word, 1).
|
||||
"""
|
||||
for i in range(1, self.maxalt[word] + 1):
|
||||
if (word,i) in self:
|
||||
for ph in self[word,i]:
|
||||
self.phoneset[ph] -= 1
|
||||
if self.phoneset[ph] == 0: # FIXME: make a class
|
||||
del self.phoneset[ph]
|
||||
dict.__delitem__(self, (word, i))
|
||||
del self.maxalt[word]
|
||||
|
||||
def alts(self, word):
|
||||
"""
|
||||
Iterate over alternative pronunciations for a word.
|
||||
"""
|
||||
for i in range(1, self.maxalt[word] + 1):
|
||||
if (word,i) in self:
|
||||
yield self[word,i]
|
||||
|
||||
@@ -9,9 +9,55 @@ class TestS3Dict(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.basedir = os.path.dirname(__file__)
|
||||
|
||||
def testCreate(self):
|
||||
def testRead(self):
|
||||
foodict = s3dict.open(os.path.join(self.basedir, "foo.dict"))
|
||||
self.assert_('AH' in foodict.phoneset)
|
||||
self.assertEquals(foodict.get_phones('A'), ['AH'])
|
||||
self.assertEquals(foodict.get_alt_phones('A', 2), ['EY'])
|
||||
self.assertEquals(foodict.get_phones('ZSWANG'), ['S', 'W', 'AE', 'NG'])
|
||||
try:
|
||||
foo = foodict.get_phones('QRXG')
|
||||
print foo
|
||||
except KeyError:
|
||||
pass # Expected fail
|
||||
else:
|
||||
self.fail()
|
||||
try:
|
||||
foo = foodict.get_alt_phones('A',3)
|
||||
except IndexError:
|
||||
pass # Expected fail
|
||||
else:
|
||||
self.fail()
|
||||
try:
|
||||
foo = foodict.get_alt_phones('!@#$!@',3)
|
||||
except KeyError:
|
||||
pass # Expected fail
|
||||
else:
|
||||
self.fail()
|
||||
self.assertEquals(foodict['A'], ['AH'])
|
||||
self.assertEquals(foodict['A',2], ['EY'])
|
||||
self.assertEquals(foodict['A(2)'], ['EY'])
|
||||
self.assertEquals(foodict['ZSWANG'], ['S', 'W', 'AE', 'NG'])
|
||||
|
||||
def testCreate(self):
|
||||
mydict = s3dict.S3Dict()
|
||||
mydict.set_phones('A', ['AH'])
|
||||
mydict.add_alt_phones('A', ['EY'])
|
||||
mydict.set_phones('ZSWANG', ['S', 'W', 'AE', 'NG'])
|
||||
mydict.set_alt_phones('A', 2, ['EY'])
|
||||
try:
|
||||
mydict.set_alt_phones('A', 5, ['AX'])
|
||||
except IndexError:
|
||||
pass # Expected fail
|
||||
else:
|
||||
self.fail()
|
||||
self.assertEquals(mydict.get_phones('A'), ['AH'])
|
||||
self.assertEquals(mydict.get_alt_phones('A', 2), ['EY'])
|
||||
self.assertEquals(mydict.get_phones('ZSWANG'), ['S', 'W', 'AE', 'NG'])
|
||||
mydict.set_alt_phones('A', 2, ['AA'])
|
||||
self.assertEquals(mydict.get_alt_phones('A', 2), ['AA'])
|
||||
mydict.del_phones('ZSWANG')
|
||||
self.assert_('NG' not in mydict.phoneset)
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user