Skip to content

Commit db33ed5

Browse files
authored
Merge pull request #431 from eric-wieser/swap-flipped-contractions
Fix incorrect right and left contraction on scalars
2 parents 48822e7 + 274591b commit db33ed5

3 files changed

Lines changed: 74 additions & 54 deletions

File tree

doc/changelog.rst

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@ Changelog
44

55
- :release:`0.5.0 <2020.05.27>`
66

7+
- :bug:`431` Left and right contraction are no longer swapped on scalar :class:`~sympy.core.expr.Expr` instances.
78
- :feature:`232` Python 3.8 is now supported and tested.
89
- :feature:`212` Python 2.7 is no longer supported, allowing more features from Python 3.x to be used.
910

galgebra/mv.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -839,7 +839,7 @@ def __lt__(self, A): # left contraction (<)
839839
)
840840

841841
if not isinstance(A, Mv): # sympy scalar
842-
return Mv(A * self.obj, ga=self.Ga)
842+
return Mv(A * self.scalar(), ga=self.Ga)
843843

844844
if self.Ga != A.Ga:
845845
raise ValueError('In < operation Mv arguments are not from same geometric algebra')
@@ -859,7 +859,7 @@ def __gt__(self, A): # right contraction (>)
859859
)
860860

861861
if not isinstance(A, Mv): # sympy scalar
862-
return self.Ga.mv(A * self.scalar())
862+
return Mv(A * self.obj, ga=self.Ga)
863863

864864
if self.Ga != A.Ga:
865865
raise ValueError('In > operation Mv arguments are not from same geometric algebra')

test/test_mv.py

Lines changed: 71 additions & 52 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,10 @@
1-
import unittest
1+
import distutils.version
22

33
import pytest
44
import sympy
55
from galgebra.ga import Ga
66

7-
class TestMv(unittest.TestCase):
7+
class TestMv:
88

99
def test_deprecations(self):
1010
ga, e_1, e_2, e_3 = Ga.build('e*1|2|3')
@@ -19,21 +19,21 @@ def test_is_base(self):
1919
"""
2020
_g3d, e_1, e_2, e_3 = Ga.build('e*1|2|3')
2121

22-
self.assertTrue((e_1).is_base())
23-
self.assertTrue((e_2).is_base())
24-
self.assertTrue((e_3).is_base())
25-
self.assertTrue((e_1 ^ e_2).is_base())
26-
self.assertTrue((e_2 ^ e_3).is_base())
27-
self.assertTrue((e_1 ^ e_3).is_base())
28-
self.assertTrue((e_1 ^ e_2 ^ e_3).is_base())
29-
30-
self.assertFalse((2 * e_1).is_base())
31-
self.assertFalse((e_1 + e_2).is_base())
32-
self.assertFalse((e_3 * 4).is_base())
33-
self.assertFalse(((3 * e_1) ^ e_2).is_base())
34-
self.assertFalse((2 * (e_2 ^ e_3)).is_base())
35-
self.assertFalse((e_3 ^ e_1).is_base())
36-
self.assertFalse((e_2 ^ e_1 ^ e_3).is_base())
22+
assert (e_1).is_base()
23+
assert (e_2).is_base()
24+
assert (e_3).is_base()
25+
assert (e_1 ^ e_2).is_base()
26+
assert (e_2 ^ e_3).is_base()
27+
assert (e_1 ^ e_3).is_base()
28+
assert (e_1 ^ e_2 ^ e_3).is_base()
29+
30+
assert not (2 * e_1).is_base()
31+
assert not (e_1 + e_2).is_base()
32+
assert not (e_3 * 4).is_base()
33+
assert not ((3 * e_1) ^ e_2).is_base()
34+
assert not (2 * (e_2 ^ e_3)).is_base()
35+
assert not (e_3 ^ e_1).is_base()
36+
assert not (e_2 ^ e_1 ^ e_3).is_base()
3737

3838
def test_get_coefs(self):
3939
_g3d, e_1, e_2, e_3 = Ga.build('e*1|2|3')
@@ -60,34 +60,34 @@ def test_blade_coefs(self):
6060
_g3d, e_1, e_2, e_3 = Ga.build('e*1|2|3')
6161

6262
m0 = 2 * e_1 + e_2 - e_3 + 3 * (e_1 ^ e_3) + (e_1 ^ e_3) + (e_2 ^ (3 * e_3))
63-
self.assertTrue(m0.blade_coefs([e_1]) == [2])
64-
self.assertTrue(m0.blade_coefs([e_2]) == [1])
65-
self.assertTrue(m0.blade_coefs([e_1, e_2]) == [2, 1])
66-
self.assertTrue(m0.blade_coefs([e_1 ^ e_3]) == [4])
67-
self.assertTrue(m0.blade_coefs([e_1 ^ e_3, e_2 ^ e_3]) == [4, 3])
68-
self.assertTrue(m0.blade_coefs([e_2 ^ e_3, e_1 ^ e_3]) == [3, 4])
69-
self.assertTrue(m0.blade_coefs([e_1, e_2 ^ e_3]) == [2, 3])
63+
assert m0.blade_coefs([e_1]) == [2]
64+
assert m0.blade_coefs([e_2]) == [1]
65+
assert m0.blade_coefs([e_1, e_2]) == [2, 1]
66+
assert m0.blade_coefs([e_1 ^ e_3]) == [4]
67+
assert m0.blade_coefs([e_1 ^ e_3, e_2 ^ e_3]) == [4, 3]
68+
assert m0.blade_coefs([e_2 ^ e_3, e_1 ^ e_3]) == [3, 4]
69+
assert m0.blade_coefs([e_1, e_2 ^ e_3]) == [2, 3]
7070

7171
a = sympy.Symbol('a')
7272
b = sympy.Symbol('b')
7373
m1 = a * e_1 + e_2 - e_3 + b * (e_1 ^ e_2)
74-
self.assertTrue(m1.blade_coefs([e_1]) == [a])
75-
self.assertTrue(m1.blade_coefs([e_2]) == [1])
76-
self.assertTrue(m1.blade_coefs([e_3]) == [-1])
77-
self.assertTrue(m1.blade_coefs([e_1, e_2, e_3]) == [a, 1, -1])
78-
self.assertTrue(m1.list() == [a, 1, -1]) # alias
74+
assert m1.blade_coefs([e_1]) == [a]
75+
assert m1.blade_coefs([e_2]) == [1]
76+
assert m1.blade_coefs([e_3]) == [-1]
77+
assert m1.blade_coefs([e_1, e_2, e_3]) == [a, 1, -1]
78+
assert m1.list() == [a, 1, -1] # alias
7979

80-
self.assertTrue(m1.blade_coefs([e_1 ^ e_2]) == [b])
81-
self.assertTrue(m1.blade_coefs([e_2 ^ e_3]) == [0])
82-
self.assertTrue(m1.blade_coefs([e_1 ^ e_3]) == [0])
83-
self.assertTrue(m1.blade_coefs([e_1 ^ e_2 ^ e_3]) == [0])
80+
assert m1.blade_coefs([e_1 ^ e_2]) == [b]
81+
assert m1.blade_coefs([e_2 ^ e_3]) == [0]
82+
assert m1.blade_coefs([e_1 ^ e_3]) == [0]
83+
assert m1.blade_coefs([e_1 ^ e_2 ^ e_3]) == [0]
8484

8585
# Invalid parameters
86-
self.assertRaises(ValueError, lambda: m1.blade_coefs([e_1 + e_2]))
87-
self.assertRaises(ValueError, lambda: m1.blade_coefs([e_2 ^ e_1]))
88-
self.assertRaises(ValueError, lambda: m1.blade_coefs([e_1, e_2 ^ e_1]))
89-
self.assertRaises(ValueError, lambda: m1.blade_coefs([a * e_1]))
90-
self.assertRaises(ValueError, lambda: m1.blade_coefs([3 * e_3]))
86+
pytest.raises(ValueError, lambda: m1.blade_coefs([e_1 + e_2]))
87+
pytest.raises(ValueError, lambda: m1.blade_coefs([e_2 ^ e_1]))
88+
pytest.raises(ValueError, lambda: m1.blade_coefs([e_1, e_2 ^ e_1]))
89+
pytest.raises(ValueError, lambda: m1.blade_coefs([a * e_1]))
90+
pytest.raises(ValueError, lambda: m1.blade_coefs([3 * e_3]))
9191

9292
def test_rep_switching(self):
9393
# this ga has a non-diagonal metric
@@ -96,29 +96,29 @@ def test_rep_switching(self):
9696
m0 = 2 * e_1 + e_2 - e_3 + 3 * (e_1 ^ e_3) + (e_1 ^ e_3) + (e_2 ^ (3 * e_3))
9797
m1 = (-4*(e_1 | e_3)-3*(e_2 | e_3))+2*e_1+e_2-e_3+4*e_1*e_3+3*e_2*e_3
9898
# m1 was chosen to make this true
99-
self.assertEqual(m0, m1)
99+
assert m0 == m1
100100

101101
# all objects start off in blade rep
102-
self.assertTrue(m0.is_blade_rep)
102+
assert m0.is_blade_rep
103103

104104
# convert to base rep
105105
m0_base = m0.base_rep()
106-
self.assertTrue(m0.is_blade_rep) # original should not change
107-
self.assertFalse(m0_base.is_blade_rep)
108-
self.assertEqual(m0, m0_base)
106+
assert m0.is_blade_rep # original should not change
107+
assert not m0_base.is_blade_rep
108+
assert m0 == m0_base
109109

110110
# convert back
111111
m0_base_blade = m0_base.blade_rep()
112-
self.assertFalse(m0_base.is_blade_rep) # original should not change
113-
self.assertTrue(m0_base_blade.is_blade_rep)
114-
self.assertEqual(m0, m0_base_blade)
112+
assert not m0_base.is_blade_rep # original should not change
113+
assert m0_base_blade.is_blade_rep
114+
assert m0 == m0_base_blade
115115

116116
def test_construction(self):
117117
ga, e_1, e_2, e_3 = Ga.build('e*1|2|3')
118118

119119
def check(x, expected_grades):
120-
self.assertEqual(x.grades, expected_grades)
121-
self.assertNotEqual(x, 0)
120+
assert x.grades == expected_grades
121+
assert x != 0
122122

123123
# non-function symbol construction
124124
check(ga.mv('A', 'scalar'), [0])
@@ -140,13 +140,13 @@ def check(x, expected_grades):
140140
check(ga.mv([1, 2, 3], 'vector'), [1])
141141

142142
# illegal arguments
143-
with self.assertRaises(TypeError):
143+
with pytest.raises(TypeError):
144144
ga.mv('A', 'vector', "too many arguments")
145-
with self.assertRaises(TypeError):
145+
with pytest.raises(TypeError):
146146
ga.mv('A', 'grade') # too few arguments
147-
with self.assertRaises(TypeError):
147+
with pytest.raises(TypeError):
148148
ga.mv('A', 'grade', not_an_argument=True) # invalid kwarg
149-
with self.assertRaises(TypeError):
149+
with pytest.raises(TypeError):
150150
ga.mv([1, 2, 3], 'vector', f=True) # can't pass f with coefficients
151151

152152
def test_abs(self):
@@ -195,3 +195,22 @@ def test_arithmetic(self):
195195
assert 1 + e_1 == one + e_1
196196
assert e_1 - 1 == e_1 - one
197197
assert 1 - e_1 == one - e_1
198+
199+
@pytest.mark.parametrize('make_one', [
200+
lambda ga: 1,
201+
lambda ga: ga.mv(sympy.S.One),
202+
pytest.param(lambda ga: sympy.S.One, marks=pytest.mark.skipif(
203+
distutils.version.LooseVersion(sympy.__version__) < "1.6",
204+
# until sympy/sympy@bec42df53cf2486d485065ddad1c31011a48bf3b
205+
reason="Cannot override < and > on sympy.Expr"
206+
))
207+
])
208+
def test_contraction(self, make_one):
209+
ga, e_1, e_2 = Ga.build('e*1|2', g=[1, 1])
210+
e12 = e_1 ^ e_2
211+
one = make_one(ga)
212+
213+
assert (one < e12) == e12
214+
assert (e12 > one) == e12
215+
assert (e12 < one) == 0
216+
assert (one > e12) == 0

0 commit comments

Comments
 (0)