Skip to content

Commit 42166d4

Browse files
author
Darren Dale
committed
more improvements to arithmetic
1 parent 77ac81a commit 42166d4

5 files changed

Lines changed: 89 additions & 55 deletions

File tree

quantities/dimensionality.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,6 @@ class BaseDimensionality(object):
1919
def simplified(self):
2020
if len(self):
2121
rq = 1*unit_registry['dimensionless']
22-
print type(rq), type(rq.dimensionality)
2322
for u, d in self.iteritems():
2423
rq *= u.reference_quantity**d
2524
return rq.dimensionality

quantities/quantity.py

Lines changed: 41 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,8 @@
1010
from quantities.registry import unit_registry
1111

1212
def prepare_compatible_units(s, o):
13+
if not isinstance(o, Quantity):
14+
o = Quantity(o, copy=False)
1315
try:
1416
assert s.dimensionality.simplified == o.dimensionality.simplified
1517
return s.simplified, o.simplified
@@ -180,7 +182,6 @@ def __add__(self, other):
180182
dims = self.dimensionality + other.dimensionality
181183
ret = super(Quantity, self).__add__(other)
182184
ret._dimensionality = dims
183-
184185
return ret
185186

186187
# TODO: in-place arithmetic should check for .base, and raise if not None
@@ -189,50 +190,42 @@ def __iadd__(self, other):
189190
if not isinstance(other, Quantity):
190191
other = Quantity(other, copy=False)
191192

192-
sd = self.dimensionality
193-
sd += other.dimensionality
194-
m = self.magnitude
195-
m += other.magnitude
196-
197-
return self
193+
self._dimensionality += other.dimensionality
194+
return super(Quantity, self).__iadd__(other)
198195

199196
def __radd__(self, other):
200197
return self.__add__(other)
201198

202199
def __sub__(self, other):
203200
if not isinstance(other, Quantity):
204-
other = Quantity(other, copy=False)
201+
other = numpy.asarray(other).view(Quantity)
205202

206203
dims = self.dimensionality - other.dimensionality
207-
magnitude = self.magnitude - other.magnitude
208-
209-
return Quantity(magnitude, dims, magnitude.dtype)
204+
ret = super(Quantity, self).__sub__(other)
205+
ret._dimensionality = dims
206+
return ret
210207

211208
def __isub__(self, other):
212209
if not isinstance(other, Quantity):
213-
other = Quantity(other, copy=False)
210+
other = numpy.asarray(other).view(Quantity)
214211

215-
sd = self.dimensionality
216-
sd -= other.dimensionality
217-
m = self.magnitude
218-
m -= other.magnitude
219-
220-
return self
212+
self._dimensionality -= other.dimensionality
213+
return super(Quantity, self).__isub__(other)
221214

222215
def __rsub__(self, other):
223216
if not isinstance(other, Quantity):
224-
other = Quantity(other, copy=False)
217+
other = numpy.asarray(other).view(Quantity)
225218

226219
dims = other.dimensionality - self.dimensionality
227-
magnitude = other.magnitude - self.magnitude
228-
229-
return Quantity(magnitude, dims, magnitude.dtype)
220+
ret = super(Quantity, self).__rsub__(other)
221+
ret._dimensionality = dims
222+
return ret
230223

231224
def __mul__(self, other):
232225
try:
233226
dims = self.dimensionality * other.dimensionality
234227
except AttributeError:
235-
other = Quantity(other, copy=False)
228+
other = numpy.asarray(other).view(Quantity)
236229
dims = Dimensionality(self.dimensionality)
237230

238231
ret = super(Quantity, self).__mul__(other)
@@ -241,14 +234,11 @@ def __mul__(self, other):
241234

242235
def __imul__(self, other):
243236
try:
244-
sd = self.dimensionality
245-
sd *= other.dimensionality
246-
m = self.magnitude
247-
m *= other.magnitude
237+
self._dimensionality *= other.dimensionality
248238
except AttributeError:
249-
m = self.magnitude
250-
m *= other
251-
return self
239+
pass
240+
241+
return super(Quantity, self).__imul__(other)
252242

253243
def __rmul__(self, other):
254244
return self.__mul__(other)
@@ -257,7 +247,7 @@ def __truediv__(self, other):
257247
try:
258248
dims = self.dimensionality / other.dimensionality
259249
except AttributeError:
260-
other = Quantity(other, copy=False)
250+
other = numpy.asarray(other).view(Quantity)
261251
dims = Dimensionality(self.dimensionality)
262252

263253
ret = super(Quantity, self).__truediv__(other)
@@ -269,21 +259,24 @@ def __div__(self, other):
269259

270260
def __itruediv__(self, other):
271261
try:
272-
sd = self.dimensionality
273-
sd /= other.dimensionality
274-
m = self.magnitude
275-
m /= other.magnitude
262+
self._dimensionality /= other.dimensionality
276263
except AttributeError:
277-
m = self.magnitude
278-
m /= other
279-
return self
264+
pass
265+
return super(Quantity, self).__itruediv__(other)
280266

281267
def __idiv__(self, other):
282268
return self.__itruediv__(other)
283269

284270
def __rtruediv__(self, other):
285-
temp = Quantity(1/self.magnitude, self.dimensionality**-1, copy=False)
286-
return other * temp
271+
try:
272+
dims = other.dimensionality / self.dimensionality
273+
except AttributeError:
274+
other = numpy.asarray(other).view(Quantity)
275+
dims = Dimensionality(self.dimensionality**-1)
276+
277+
ret = super(Quantity, self).__rtruediv__(other)
278+
ret._dimensionality = dims
279+
return ret
287280

288281
def __rdiv__(self, other):
289282
return self.__rtruediv__(other)
@@ -294,41 +287,39 @@ def __pow__(self, other):
294287
raise ValueError("exponent must be dimensionless")
295288
other = other.simplified.magnitude
296289

297-
other = numpy.array(other)
290+
other = numpy.asarray(other)
298291
try:
299292
assert other.min() == other.max()
300293
other = other.min()
301294
except AssertionError:
302295
raise ValueError('Quantities must be raised to a single power')
303296

304297
dims = self.dimensionality**other
305-
magnitude = self.magnitude**other
306-
return Quantity(magnitude, dims, magnitude.dtype)
298+
ret = super(Quantity, self).__pow__(other)
299+
ret._dimensionality = dims
300+
return ret
307301

308302
def __ipow__(self, other):
309303
if isinstance(other, Quantity):
310304
if other.dimensionality.simplified:
311305
raise ValueError("exponent must be dimensionless")
312306
other = other.simplified.magnitude
313307

314-
other = numpy.array(other)
308+
other = numpy.asarray(other)
315309
try:
316310
assert other.min() == other.max()
317311
other = other.min()
318312
except AssertionError:
319313
raise ValueError('Quantities must be raised to a single power')
320314

321-
sd = self.dimensionality
322-
sd **= other
323-
m = self.magnitude
324-
m **= other
325-
return self
315+
self._dimensionality **= other
316+
return super(Quantity, self).__ipow__(other)
326317

327318
def __rpow__(self, other):
328319
if self.dimensionality.simplified:
329320
raise ValueError("exponent must be dimensionless")
330321

331-
return other**self.simplified.magnitude
322+
return super(Quantity, self.simplified).__rpow__(other)
332323

333324
def __repr__(self):
334325
return '%s %s'%(numpy.ndarray.__str__(self), self.dimensionality)

quantities/tests/test_quantities.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -494,7 +494,7 @@ def test_multiplication(self):
494494
103.0 * q.kPa*q.inch
495495
)
496496

497-
self.assertAlmostEqual((5.2 * q.eV) * (300.2 * q.eV), 1561.04 * q.eV**2)
497+
self.assertAlmostEqual((5.2 * q.J) * (300.2 * q.J), 1561.04 * q.J**2)
498498

499499
# the formatting should be the same
500500
self.assertEqual(

quantities/uncertainquantity.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -137,10 +137,10 @@ def __rtruediv__(self, other):
137137
return other * temp
138138

139139
def __pow__(self, other):
140-
res = Quantity.__pow__(self, other)
140+
res = super(UncertainQuantity, self).__pow__(other)
141141
ru = other * self.relative_uncertainty
142-
u = res * ru
143-
return UncertainQuantity(res, uncertainty=u, copy=False)
142+
res.uncertainty = res * ru
143+
# return UncertainQuantity(res, uncertainty=u, copy=False)
144144

145145
def __getitem__(self, key):
146146
return UncertainQuantity(

quantities/unitquantity.py

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,17 @@
1414
]
1515

1616

17+
def quantity(f):
18+
19+
def wrapped(*args, **kwargs):
20+
ret = f(*args, **kwargs)
21+
if isinstance(ret, UnitQuantity):
22+
return ret.view(Quantity).copy()
23+
return ret
24+
25+
return wrapped
26+
27+
1728
class UnitQuantity(Quantity):
1829

1930
_primary_order = 99
@@ -73,6 +84,39 @@ def __repr__(self):
7384
return s+'\nnote: %s'%self.note
7485
return s
7586

87+
@quantity
88+
def __mul__(self, other):
89+
return super(UnitQuantity, self).__mul__(other)
90+
91+
@quantity
92+
def __rmul__(self, other):
93+
return super(UnitQuantity, self).__rmul__(other)
94+
95+
@quantity
96+
def __truediv__(self, other):
97+
return super(UnitQuantity, self).__truediv__(other)
98+
99+
@quantity
100+
def __rtruediv__(self, other):
101+
return super(UnitQuantity, self).__rtruediv__(other)
102+
103+
@quantity
104+
def __div__(self, other):
105+
return super(UnitQuantity, self).__div__(other)
106+
@quantity
107+
def __rdiv__(self, other):
108+
return super(UnitQuantity, self).__rdiv__(other)
109+
110+
def __imul__(self, other):
111+
raise TypeError('can not modify protected units')
112+
113+
def __itruediv__(self, other):
114+
raise TypeError('can not modify protected units')
115+
116+
def __idiv__(self, other):
117+
raise TypeError('can not modify protected units')
118+
119+
76120
@property
77121
def format_order(self):
78122
return self._format_order

0 commit comments

Comments
 (0)