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:
dhdfu
2010-03-07 02:53:52 +00:00
parent 69bf793aed
commit 2db07abe32
2 changed files with 185 additions and 31 deletions
+138 -30
View File
@@ -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]
+47 -1
View File
@@ -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()